Skip to content

feat(libtorch): piecewise-linear table lookup composed from float32 primitives - #34

Open
NicolasRouquette wants to merge 1 commit into
lean-dojo:mainfrom
NicolasRouquette:libtorch-table-lookup
Open

NicolasRouquette wants to merge 1 commit into
lean-dojo:mainfrom
NicolasRouquette:libtorch-table-lookup

Conversation

@NicolasRouquette

Copy link
Copy Markdown
Contributor

What

A tabulated function (a calibration curve, an inverted sensor response, an activation lookup)
evaluated at grid coordinates by clamp-addressed linear interpolation: the uniform-grid analogue
of np.interp, as a CUDA tape node with gradients for both the table and the coordinates.

  • Tape.tableLookup {width} tableId coordsId for a rank-one table of width samples, and
    Tape.layeredTableLookup {layers width} tableId coordsId layerId for a layers × width table
    whose layer is selected per element from a third tape value. The result has the coordinates'
    shape. With i = ⌊u⌋ and w = u - i after clamping u to [0, width - 1], the value is
    v₀ + (v₁ - v₀) · w with v₀ = table[i] and v₁ = table[min (i + 1) (width - 1)].
  • Two primitives that did not exist and that the node is composed from: floor in the shared
    operations.h list (so the unavailable build gets its twin for free), and Buffer.gatherAt
    with its VJP Buffer.scatterAddAt, a gather at positions that live on the device as a float32
    buffer. Both also get tape nodes (Tape.floor, Tape.gatherAt).
  • KernelSpec.tableLookupSpec, the ExecFloat.Binary 8 23 reference the forward matches.

Relation to #24

#24 put the same function behind a layered cudaArray and a texture object, which is the kind of
custom CUDA this backend no longer has. The function survives the move; the mechanism does not.
What this PR keeps from #24 is the contract that mattered there: the forward is bit-identical to
an executable Float32 reference, because every operation it is composed from rounds exactly
once (each ATen elementwise kernel is one rounding, and nothing contracts across kernel launches).
What it adds is the backward #24 deferred: the coordinate receives the slope of its segment, zero
where the clamp is active (the same open-interval rule as clamp), and the table receives the two
interpolation weights scattered to the samples they multiplied. The 9-bit hardware filter mode
is not reproducible on ATen and is not offered.

Design notes

  • Positions on device. gatherAt takes positions as a float32 buffer, which is what the tape
    computes: a NaN position reads as 0, every position is clamped to [0, size - 1] and then
    truncated, so a non-integral position floors and no selection leaves the table without a
    device-side assertion. The tape operation rejects empty tables and tables above 2^24
    samples, which is where float32 positions stop being exact.
  • One rounding per operation. tableSegments computes the clamped coordinate, its floor, the
    weight, and the neighbouring position with clamp, floor, sub, add; all of these are
    exact in float32 below 2^24. Only the final sub, mul, add round, in that order, which is
    the order tableLookupSpec states.
  • Ownership. The three position buffers are kept on the node's cleanup list for the
    backward, so backward does not recompute them, and the eager session releases them with the
    tape. Every other intermediate is released through the value that consumed it.
  • Cost. A rank-one lookup is about a dozen small ATen calls; the layered one a few more. This
    is a correctness and composability change, not a performance claim, and no timing is reported.
  • Accumulation order. scatterAddAt accumulates repeated positions through index_add_,
    whose order is the SDK's. A repeatability check under the deterministic control is included
    with the other deterministic reductions; the gradient comparison uses dyadic values, where the
    order does not change the sum.

Tests

NN/Tests/Runtime/Cuda/TableLookup.lean, wired into the CUDA coverage suite:

  • floor agrees with Float32.floor bit for bit, including signed zeros, 16777217, and
    ±1e8; its tape node reports a zero gradient.
  • gatherAt agrees with the CPU tape's indexSelect exactly, forward and table gradient,
    including a repeated position; and the documented addressing holds for negative, past-the-end,
    non-integral, and NaN positions.
  • Both lookups match the Float32 reference bit for bit over interior, boundary, out-of-range,
    integral, signed-zero, and NaN coordinates, with layer selectors past the end and non-integral.
  • The forward and both gradients agree with the same function recorded on the CPU tape from
    existing nodes (clamp, sub, indexSelect twice, sub, mul, add, with the segment
    indices computed on the host by the same recipe).
  • Edges: a one-sample table (every coordinate reads it, no coordinate gradient, the table
    gradient is the seed sum), zero coordinates, and the two rejections.

NN/Tests/Runtime/Cuda/DeterministicReductions.lean gains scatterAddAt repeatability under the
deterministic control, next to the existing scatterAdd check.

Checks run

RTX A4500 (sm_86, driver 580.126), pip torch 2.11.0+cu128 as the SDK, CUDA 12.8 toolkit for
SDK discovery, Lean v4.34.0:

scripts/libtorch_build.py (torchlean.cpp, C++20, GNU 11.5)    Built target torchlean_libtorch
lake -R -Kcuda=true build nn_tests_suite                       Build completed successfully
TORCHLEAN_REQUIRE_CUDA=1 nn_tests_suite                        === CUDA kernel coverage: floor, gather_at, table lookup ===
                                                               == CUDA deterministic reductions: OK ==
                                                               == TorchLean: all curated tests passed ==
lake -R -Kcuda=false build NN NNTests nn_tests_suite NNCI NNExamples   Build completed successfully
nn_tests_suite (default build)                                 CUDA kernels: skipped (LibTorch not linked)
                                                               == TorchLean: all curated tests passed ==
python3 scripts/checks/repo_lint.py                            OK: no issues found.

Not run: the sanitizer harness and the elementwise C++ harness.

…rimitives

Add `floor` to the shared operation list, a gather at device-resident positions with its
scatter-add VJP, and `Tape.tableLookup` / `Tape.layeredTableLookup`: clamp-addressed linear
interpolation into a tabulated function, composed in Lean from `clamp`, `floor`, `sub`, `mul`,
`add`, and the new gather so that every operation rounds once and the result matches
`KernelSpec.tableLookupSpec` bit for bit. The backward gives the coordinate its segment's slope,
zero where the clamp is active, and the table the two interpolation weights.

Tests compare `floor` with `Float32.floor`, `gatherAt` with the CPU tape's `indexSelect`, the
lookups with a `Float32` reference at boundary, out-of-range, integral, and NaN coordinates, and
the forward and both gradients with the same function recorded on the CPU tape.
NicolasRouquette added a commit to NicolasRouquette/TorchLean that referenced this pull request Oct 3, 2026
Brings in the upstream PR branch (lean-dojo#34) at 558a3af, one
commit on upstream main b062b9a: floor in operations.h, Buffer.gatherAt
with its scatterAddAt VJP, Tape.floor, Tape.gatherAt, and
Tape.tableLookup / Tape.layeredTableLookup composed from float32 primitives
so the forward matches KernelSpec.tableLookupSpec bit for bit.

Successor of the closed lean-dojo#24 (TexTable). The texture mechanism and the
hardware filter mode do not carry over to this backend.
@Robertboy18

Copy link
Copy Markdown
Member

Hey Nicolas, thanks for the interpolation work. I found a bounds issue in table_positions: the upper clamp bound is converted to float32 before the integer conversion. With size = 16777220 and float32 positions [inf, 16777220], the same ATen nan_to_num/clamp/to(int64) sequence produces indices [16777220, 16777220], but the last valid index is 16777219. tableLookup limits its table size, but the exposed gatherAt/scatterAddAt primitives do not have that limit. Could you clamp again in the integer domain to [0, size - 1] and add a regression above 2^24, including infinity? Holding off on merging until this is fixed.

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.

2 participants