Repository navigation
feat(libtorch): piecewise-linear table lookup composed from float32 primitives - #34
Open
NicolasRouquette wants to merge 1 commit into
Open
NicolasRouquette wants to merge 1 commit into
NicolasRouquette wants to merge 1 commit into
Conversation
…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.
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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 coordsIdfor a rank-one table ofwidthsamples, andTape.layeredTableLookup {layers width} tableId coordsId layerIdfor alayers × widthtablewhose layer is selected per element from a third tape value. The result has the coordinates'
shape. With
i = ⌊u⌋andw = u - iafter clampinguto[0, width - 1], the value isv₀ + (v₁ - v₀) · wwithv₀ = table[i]andv₁ = table[min (i + 1) (width - 1)].floorin the sharedoperations.hlist (so the unavailable build gets its twin for free), andBuffer.gatherAtwith its VJP
Buffer.scatterAddAt, a gather at positions that live on the device as a float32buffer. Both also get tape nodes (
Tape.floor,Tape.gatherAt).KernelSpec.tableLookupSpec, theExecFloat.Binary 8 23reference the forward matches.Relation to #24
#24 put the same function behind a layered
cudaArrayand a texture object, which is the kind ofcustom 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
Float32reference, because every operation it is composed from rounds exactlyonce (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 twointerpolation weights scattered to the samples they multiplied. The 9-bit hardware filter mode
is not reproducible on ATen and is not offered.
Design notes
gatherAttakes positions as a float32 buffer, which is what the tapecomputes: a NaN position reads as
0, every position is clamped to[0, size - 1]and thentruncated, 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^24samples, which is where float32 positions stop being exact.
tableSegmentscomputes the clamped coordinate, its floor, theweight, and the neighbouring position with
clamp,floor,sub,add; all of these areexact in float32 below
2^24. Only the finalsub,mul,addround, in that order, which isthe order
tableLookupSpecstates.cleanuplist for thebackward, 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.
is a correctness and composability change, not a performance claim, and no timing is reported.
scatterAddAtaccumulates repeated positions throughindex_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:flooragrees withFloat32.floorbit for bit, including signed zeros,16777217, and±1e8; its tape node reports a zero gradient.gatherAtagrees with the CPU tape'sindexSelectexactly, forward and table gradient,including a repeated position; and the documented addressing holds for negative, past-the-end,
non-integral, and NaN positions.
Float32reference bit for bit over interior, boundary, out-of-range,integral, signed-zero, and NaN coordinates, with layer selectors past the end and non-integral.
existing nodes (
clamp,sub,indexSelecttwice,sub,mul,add, with the segmentindices computed on the host by the same recipe).
gradient is the seed sum), zero coordinates, and the two rejections.
NN/Tests/Runtime/Cuda/DeterministicReductions.leangainsscatterAddAtrepeatability under thedeterministic control, next to the existing
scatterAddcheck.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:
Not run: the sanitizer harness and the elementwise C++ harness.