Skip to content

fix(dag): infer argMax output spec as reduced i32 (fixes #876) - #878

Merged
michalharakal merged 1 commit into
developfrom
fix/argmax-dag-output-spec
Jul 24, 2026
Merged

michalharakal merged 1 commit into
developfrom
fix/argmax-dag-output-spec

Conversation

@michalharakal

Copy link
Copy Markdown
Contributor

The dag{} builder path never inferred argMax's output spec. GraphDsl.inferDagOutputSpecs has cases for
reductions / 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)
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 board-verified and only the raw dag{} path broke — surfaced by the functiongemma-270m
conformance row (skainet-iree-conformance #24).

Fix

Add an argmax/argmin case to inferDagOutputSpecs — reduced shape (like sum/mean) with an Int32
index dtype — and let the local spec() helper 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 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-hlo argMax tests 5/5, full skainet-lang-dag suite green.

Fixes #876.

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>
@github-actions

Copy link
Copy Markdown

📖 Documentation Preview

The documentation has been built successfully for this PR.

Generated Files:

  • Operator documentation: docs/modules/operators/_generated_/
  • JSON schema output: operators.json

Artifacts:

  • Download the documentation-preview-878 artifact to view the complete documentation locally.

This comment will be updated automatically when the PR is updated.

@michalharakal
michalharakal requested a review from aharakal July 24, 2026 17:17
@michalharakal
michalharakal merged commit b7e4ed4 into develop Jul 24, 2026
15 of 20 checks passed
@michalharakal
michalharakal deleted the fix/argmax-dag-output-spec branch July 24, 2026 17:20
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>
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.

argMax → StableHLO lowering emits invalid IR (int-literal f32 constant + non-collapsing reduce)

2 participants