fix(dag): infer argMax output spec as reduced i32 (fixes #876) - #878
Merged
Merged
Conversation
The `dag{}` builder path never inferred argMax's output spec: `GraphDsl.inferDagOutputSpecs`
had cases for reductions/reshape/matmul/concat but not argMax, so it fell back to echoing
operand-0's shape + dtype. An argMax over f32 logits therefore recorded an f32, full-shape
output; the (correct) StableHloConverter then faithfully emitted invalid IR:
- `stablehlo.constant dense<N> : tensor<...xf32>` — an integer literal for an f32 tensor,
which iree-compile rejects ("unexpected decimal integer literal for a floating point value");
- a final `stablehlo.reduce` whose result kept the reduced dim (`1x4x8` not `1x4`).
The VoidTensorOps path already inferred this correctly (reduced + Int32), which is why the
FunctionGemma NN-DSL export was fine and only the raw `dag{}` path broke — surfaced by the
functiongemma-270m conformance row.
Fix: add an `argmax`/`argmin` case to `inferDagOutputSpecs` — reduced shape (like sum/mean)
with an `Int32` index dtype — and let `spec()` take a dtype override. No converter change; it
was correct given a correct spec.
Test: `ArgMaxOperationsConverterTest.testArgMaxThroughDagBuilderProducesCompilableIr` traces
`dag { argMax }` end to end and asserts i32 sentinel + collapsed reduce + no int-literal-f32
constant — the exact conditions #876 broke. (The pre-existing converter test used a hand-built
correct node, bypassing the inference — which is why the bug slipped.)
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
|
📖 Documentation Preview The documentation has been built successfully for this PR. Generated Files:
Artifacts:
This comment will be updated automatically when the PR is updated. |
aharakal
approved these changes
Jul 24, 2026
MacOS
pushed a commit
to MacOS/SKaiNET
that referenced
this pull request
Jul 27, 2026
Bump version 0.36.0 -> 0.37.0 (gradle.properties, docs/antora.yml, README quickstart). Promote CHANGELOG [Unreleased] to [0.37.0]: Lstm layer (SKaiNET-developers#824), real Dropout masking (SKaiNET-developers#867), LR schedules and mutable optimizer lr (SKaiNET-developers#866), optional-bias and open Linear (SKaiNET-developers#870, SKaiNET-developers#875), androidNative IO targets (SKaiNET-developers#836, SKaiNET-developers#842, SKaiNET-developers#845), the SDPA default-scale fix (SKaiNET-developers#880), three autograd fixes (SKaiNET-developers#877), the argMax DAG output spec (SKaiNET-developers#878), tokenizer BPE inference and N-D gather (SKaiNET-developers#879), plus the CI/docs supply-chain hardening and toolchain bumps. Refresh README "What's New" and add a Contributors (0.37.0) section. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
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.
The
dag{}builder path never inferredargMax's output spec.GraphDsl.inferDagOutputSpecshas cases forreductions / reshape / matmul / concat / conv but not argMax, so it fell back to echoing operand-0's
shape + dtype. An argMax over f32 logits thus recorded an f32, full-shape output; the (correct)
StableHloConverterthen faithfully emitted invalid IR:stablehlo.constant dense<N> : tensor<…xf32>— an integer literal for an f32 tensor, whichiree-compilerejects ("unexpected decimal integer literal for a floating point value");stablehlo.reducewhose result kept the reduced dim (1x4x8, not1x4).The
VoidTensorOpspath already inferred this correctly (reduced +Int32), which is why the FunctionGemmaNN-DSL export was board-verified and only the raw
dag{}path broke — surfaced by thefunctiongemma-270mconformance row (skainet-iree-conformance #24).
Fix
Add an
argmax/argmincase toinferDagOutputSpecs— reduced shape (likesum/mean) with anInt32index dtype — and let the local
spec()helper take a dtype override. No converter change — it wascorrect given a correct spec.
Test
ArgMaxOperationsConverterTest.testArgMaxThroughDagBuilderProducesCompilableIrtracesdag { argMax }end toend and asserts the i32 sentinel + collapsed reduce + no int-literal-f32 constant — the exact conditions #876
broke. The pre-existing converter test used a hand-built correct node, bypassing the inference (why the bug
slipped). Verified:
skainet-compile-hloargMax tests 5/5, fullskainet-lang-dagsuite green.Fixes #876.