Skip to content

feat: GIN, a cascade of random convolutions for domain generalisation - #131

Open
Hendrik-code wants to merge 1 commit into
mainfrom
hm/gin
Open

Hendrik-code wants to merge 1 commit into
mainfrom
hm/gin

Conversation

@Hendrik-code

Copy link
Copy Markdown
Collaborator

What does this change?

Adds RandomGINGPU, a GPU implementation of GIN (Global Intensity Non-linear
augmentation): the image is pushed through a shallow convolutional network whose weights
are drawn afresh on every call, blended back towards the original at a per-sample alpha,
and rescaled so the output carries the same energy as the input.

Ouyang, C., Chen, C., Li, S., Li, Z., Qin, C., Bai, W., & Rueckert, D. (2022).
Causality-inspired single-source domain generalization for medical image segmentation.
IEEE TMI 42(4), 1095-1106. DOI 10.1109/TMI.2022.3224067

Ported from the authors' own 3-D implementation (models/imagefilter3d.py in
cheng-01037/Causality-Medical-Image-Domain-Generalization)
and the cleaner nnU-Net-oriented restatement in dg_tta/gin.py
(multimodallearning/DG-TTA). Both are MIT
licensed; this is a fresh implementation in SmaugLab's idiom with attribution in the module
docstring.

New files: smauglab/transforms/gpu/gin.py, unit_tests/test_gin.py. Wiring:
AugId.GIN + a PIPELINE_ORDER slot in registry.py, the import in
transforms/__init__.py, the module in TRANSFORM_MODULES, and the regenerated
README.md matrix row and all_augmentations.json block.

Why?

RandomRandConvGPU was the closest thing we had and is single-layer RandConv: one
random kernel, no non-linearity, no output renormalisation. GIN is its multi-layer
non-linear generalisation, and the one structural idea the transfer-augmentation set was
missing. It gets its own AugId rather than sharing RAND_CONV — sharing would have
overwritten RandomRandConvGPU's matrix cell.

How was it tested?

  • pytest passes locally — 472 passed, 15 skipped, 1665 subtests (the skips are the
    pre-existing $SMAUGLAB_DOMAIN_BANK and build-package ones)
  • pre-commit run --all-files passes
  • Added tests covering the change — unit_tests/test_gin.py, 19 tests / 64 subtests
  • Ran a real GPU batch end to end: [2, 1, 128, 128, 128] on one A40 inside
    torch.autocast("cuda"), output finite, energy-matched, 3.8–5.5 ms per batch
    against 0.7 ms for single-layer RandConv

mypy smauglab/ is clean. Note my local mypy is 2.1.0 and ruff 0.15.20, not the pinned
2.4.0 / 0.16.1 — but pre-commit installs its own pinned ruff and that passed, so only the
mypy version is unverified against CI's.

Anything reviewers should look at closely?

The float16 guard is the part I would check first. The trainer runs the GPU pipeline
inside autocast, so F.conv3d would hand back float16; the kernels are raw N(0, 1) with
no fan-in scaling, and simulating the stack without the forcing, ||mixed||_F passed
float16's 65504 in 23 of 150 draws at the defaults on a 1×128³ patch, 48 of 150 on
2×192³, and 105 of 150 with kernel_sizes=(1, 3, 5, 7). The failure is not NaN, and that
is the whole problem:
the norm comes back inf, 1 / (inf + 1e-5) is exactly 0., and
the output is an all-zero patch that torch.isfinite(...).all() reports as fine — so
neither _select_and_check nor the trainer's non-finite-loss guard can see it, and the
network trains on a blank volume. apply_transform therefore runs the cascade and the norm
inside an autocast_active()-guarded enabled=False region, with promote_types rather
than .float() so a float64 input keeps its precision.
test_a_float16_input_stays_finite_and_energy_matched uses the configuration measured to
fail in 19 of 20 draws without it, and asserts the energy match rather than just finiteness.

Two deliberate departures from upstream, both argued at their call sites:

  1. Reflect padding rather than zeros. Every other convolution in this package
    reflect-pads, and a zero-padded 3×3×3 kernel darkens the patch border, which an nnU-Net
    patch has no business having — it is an interior crop, not a whole image.
    padding_mode="zeros" restores upstream exactly.
  2. float32 as a floor, above.

No fan-in scaling on the kernels, unlike _RandomConvBaseGPU.get_kernel's
1/sqrt(k**3). That is safe here and not there, because the Frobenius renormalisation is
exact rather than statistical: the per-sample output RMS matches the input's to within
4.1e-07 relative for every combination of n_layers ∈ {1,2,4,8} and
kernel_sizes ∈ {(1,), (3,), (1,3), (1,3,5,7)}. The invariant is on RMS, not on the
standard deviation — the biases and the leaky_relu move energy into the DC term, so the
standard deviation still varies by about 1.3x across those settings, which is why this does
not copy test_randconv_gain.py's metric.
test_without_the_frobenius_norm_the_scale_does_track_the_configuration is the control:
out_norm="none" lets the scale run away by 19197x over the same sweep.

Joint channel mapping. apply_to_channel is mapped by one network with that many
inputs and outputs, because cross-channel mixing is the published method. For
single-channel data — the usual nnU-Net case — that is identical to a per-channel loop. It
does mean a duplicate entry changes the network's width rather than merely applying it
twice, so (0, 0) is rejected at construction; every sibling treats it as harmless.

No mix_prob, unlike the siblings: alpha_range already is the blend back towards
the original, and a mix_prob on top would blend twice.

Checklist

  • Augmentation behaviour is unchanged, or the change is intentional and described above

    Unchanged. transform_params_gpu.json gains "RandomGINGPU": {"p": 0.0} so the default
    config names it and a sweep only has to flip one number. The config has no pipeline
    section so it runs SEQUENTIAL, and _leaf_parameters short-circuits p == 0 to
    torch.zeros(batch) without touching the RNG. Verified by building the pipeline from the
    config before and after under one seed: bitwise identical, max abs diff 0.0, against
    2.886 for the same entry at p=1.0. Only the content hash moves, 49138fe4 → 1c7447e7.
    No other transform_params_*.json is touched.

  • New transforms are reachable from a config JSON

    Via the regenerated all_augmentations.json and via transform_params_gpu.json above.

🤖 Generated with Claude Code

GIN pushes the image through a shallow convolutional network whose weights are
drawn afresh on every call, blends the result back towards the original at a
per-sample alpha, and rescales so the output carries the same energy as the
input. Because the network is random rather than trained, the family of
intensity mappings it realises is far wider than any fixed curve, which is the
point: it is a single-source domain-generalisation augmentation.

    Ouyang, C., Chen, C., Li, S., Li, Z., Qin, C., Bai, W., & Rueckert, D.
    (2022). Causality-inspired single-source domain generalization for medical
    image segmentation. IEEE TMI 42(4), 1095-1106. DOI 10.1109/TMI.2022.3224067

Ported from the authors' own 3-D implementation (`models/imagefilter3d.py` in
cheng-01037/Causality-Medical-Image-Domain-Generalization) and the cleaner
nnU-Net-oriented restatement in `dg_tta/gin.py`
(multimodallearning/DG-TTA). Both are MIT licensed.

`RandomRandConvGPU` was the closest thing here and is single-layer RandConv:
one random kernel, no non-linearity, no output renormalisation. GIN is its
multi-layer non-linear generalisation, so it gets its own AugId -- sharing
RAND_CONV's would have overwritten that transform's matrix cell -- and sits
next to it in PIPELINE_ORDER.

Two deliberate departures from upstream, both argued at their call sites:

* Reflect padding rather than zeros. Every other convolution in this package
  reflect-pads, and a zero-padded 3x3x3 kernel darkens the patch border, which
  an nnU-Net patch has no business having: it is an interior crop, not a whole
  image. `padding_mode="zeros"` restores upstream exactly, and
  `test_zero_padding_changes_only_the_border` pins that the choice reaches
  exactly as far into the patch as the cascade can.

* float32 as a floor for the cascade and the norm. The trainer runs the GPU
  pipeline inside `autocast`, so `F.conv3d` would hand back float16; the
  kernels are raw N(0, 1) with no fan-in scaling, and simulating the stack
  without the forcing, `||mixed||_F` passed float16's 65504 in 23 of 150 draws
  at the defaults on a 1x128^3 patch, 48 of 150 on 2x192^3, and 105 of 150 with
  kernel_sizes=(1, 3, 5, 7). The failure is not NaN and that is the whole
  problem: the norm comes back `inf`, `1 / (inf + 1e-5)` is exactly 0, and the
  output is an all-zero patch that `torch.isfinite(...).all()` reports as fine.
  Neither `_select_and_check` nor the trainer's non-finite-loss guard can see
  it, and the network trains on a blank volume. `promote_types` rather than
  `.float()` so a float64 input keeps its precision.

The unscaled kernels are safe because of the Frobenius renormalisation, which
is exact rather than statistical: the per-sample output RMS matches the input's
to within 4.1e-07 relative for every combination of n_layers in {1, 2, 4, 8}
and kernel_sizes in {(1,), (3,), (1,3), (1,3,5,7)}. That is why GIN needs no
`1/sqrt(k**3)` where RandConv does. The invariant is on RMS and not on the
standard deviation: the biases and the leaky_relu move energy into the DC term,
so the standard deviation still varies by about 1.3x across those settings.
`test_without_the_frobenius_norm_the_scale_does_track_the_configuration` is the
control -- turning the renormalisation off lets the scale run away by 19197x
over the same sweep.

The channels named by `apply_to_channel` are mapped jointly, by one network
with that many inputs and outputs, because cross-channel mixing is the
published method. For single-channel data, which is the usual nnU-Net case,
that is identical to a per-channel loop. It does make a duplicate entry change
the network's width rather than merely applying it twice, so `(0, 0)` is now
rejected at construction.

There is no `mix_prob`, unlike the sibling transforms: `alpha_range` already is
the blend back towards the original, and a `mix_prob` on top would blend twice.

Augmentation behaviour is unchanged. `transform_params_gpu.json` gains
`"RandomGINGPU": {"p": 0.0}` so the default config names it and a sweep only
has to flip one number; the config has no `pipeline` section so it runs
SEQUENTIAL, and `_leaf_parameters` short-circuits p == 0 without touching the
RNG. Verified by building the pipeline from the config before and after under
one seed: bitwise identical, max abs diff 0.0, against 2.886 for the same entry
at p=1.0. Only the content hash moves, 49138fe4 -> 1c7447e7.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant