Skip to content

Enforce validation when creating an MLOperandDescriptor - #955

Open
shiyi9801 wants to merge 3 commits into
webmachinelearning:mainfrom
shiyi9801:validate_lstm_gru
Open

shiyi9801 wants to merge 3 commits into
webmachinelearning:mainfrom
shiyi9801:validate_lstm_gru

Conversation

@shiyi9801

@shiyi9801 shiyi9801 commented Sep 11, 2026 •

Copy link
Copy Markdown
Contributor

Fix #949

This PR makes create an MLOperandDescriptor validate the resulting shape via checking dimensions, so oversized outputs are rejected during graph building. concat(), pad(), and the pooling ops are restructured to construct their descriptor from the final computed shape instead of mutating one created earlier.

This PR also adds unconditional validation for the full-sequence output of gru()/lstm(), which some platforms produce regardless of the returnSequence flag.


Preview | Diff

Comment thread index.bs
1. Let |descriptor| be a new {{MLOperandDescriptor}}.
1. Set |descriptor|.{{MLOperandDescriptor/dataType}} to |dataType|.
1. Set |descriptor|.{{MLOperandDescriptor/shape}} to a [=list/clone=] of |shape|.
1. If [=MLOperandDescriptor/checking dimensions=] given |descriptor| returns false, then [=exception/throw=] a {{TypeError}}.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is a good fix — it ensures all output operand descriptors are validated, which was missing before.

Since this now applies to every computed output descriptor (not just GRU/LSTM), could you mention that scope in the PR title and commit log? A number of operators — conv2d, gemm, matmul, gather, tile, expand, where, quantize/dequantizeLinear, argMin/argMax, and the broadcasting element-wise ops — gain new TypeError paths here, and reviewers should know that going in.

Two gaps the constructor-side check can't reach:

concat constructs the descriptor from first's shape and then mutates shape[axis] to the accumulated size, so the validated shape isn't the final one. It needs a check after the mutation.

reshape copies input's descriptor and sets shape directly without calling the constructor. The element count is already covered by the existing product check, but the byte length / rank limits aren't.

Could you also audit for any other methods that bypass the constructor or change dimensions afterwards?

@shiyi9801 shiyi9801 Sep 14, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could you also audit for any other methods that bypass the constructor or change dimensions afterwards?

Done! Fixed the constructor of concat(), pad(), and the pooling ops.

reshape already validates that the output tensor has the same element count as the input tensor, which is ensured to be valid, so it does not need additonal checks.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks — concat looks right now, and good catch on pad and pool2d.

reshape still copies input.[[descriptor]] and sets shape directly without going through the constructor, so it doesn't get the new validation. As you said, the element count is already covered by the existing product check, but the rank limit isn't — could you convert it to use creating an MLOperandDescriptor for consistency?

@shiyi9801 shiyi9801 changed the title Validate LSTM/GRU full-sequence output regardless of the returnSequence flag Enforce validation when creating an MLOperandDescriptor Sep 14, 2026
@anssiko

anssiko commented Sep 28, 2026

Copy link
Copy Markdown
Member

@shiyi9801 just confirming the shape of this PR is fine and is covered by the existing resolution.

After the PR review comments are addressed, this can be merged.

Thank you for your contributions that make the API more secure.

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.

Validate LSTM/GRU full-sequence output regardless of the returnSequence flag

3 participants