siyiweigeHEW opened a new pull request, #20149:
URL: https://github.com/apache/tvm/pull/20149

   
   Fixes: #20148
   
   ## Summary
   
   The Relax ONNX frontend `PRelu._impl_v1` rejected a valid `PRelu` whose 
`slope`
   is broadcastable to `X` across **multiple** non-broadcast axes (e.g. a slope
   shaped exactly like `X`, or with several non-1 dims) with a `ValueError`, 
even
   though onnxruntime and `onnx.reference` accept and run such models correctly.
   After #20115 added lower-rank and rank-0 slope support, this was the 
remaining
   gap for unidirectionally broadcastable slopes.
   
   ## Root cause
   
   `PRelu._impl_v1` lowers a slope with a **single** non-broadcast axis to
   `relax.op.nn.prelu(x, slope_vec, axis)`, which is the only shape Relax's
   `nn.prelu` can express (one per-axis slope vector). A slope that broadcasts
   along several axes (`slope == X`, `(1, C, H, W)`, `(N, C, 1, 1)`, …) cannot 
be
   represented that way, so the code raised instead of lowering the op:
   
   ```python
   # Must have only ONE non-broadcast axis
   if len(non_one_axes) != 1:
       raise ValueError(
           f"Invalid PRelu slope shape (multiple non-broadcast dims): 
{slope_shape}"
       )
   ```
   
   ## Fix
   
   For the multi-non-broadcast case, lower `PRelu` elementwise instead:
   
   ```python
   dtype = x.ty.dtype.dtype
   return relax.op.where(
       relax.op.less(x, relax.const(0, dtype)),
       relax.op.multiply(x, slope),
       x,
   )
   ```
   
   This is exactly `PRelu(x, s) = where(x < 0, s * x, x)` with `s` broadcast to
   `x`. The all-ones / rank-1 / single-non-broadcast paths are unchanged and 
still
   use `nn.prelu`. Non-broadcastable slopes (rank > x, or a non-1 dim that does 
not
   match `x`) still fail at build time with a broadcast error, consistent with
   onnxruntime rejecting them at session creation.
   
   ## Validation
   
   Differential test: Relax (build + `VirtualMachine`) vs onnxruntime on the 
same
   model, comparing output shapes and max |diff| (cross-checked with
   `onnx.reference`), across 7 `X` shapes × every valid broadcastable `slope`.
   
   | Case | onnxruntime | TVM | max\|diff\| | Result |
   |---|---|---|---|---|
   | lower-rank `(64,1,1)` on `(1,64,128,128)` | `(1,64,128,128)` | 
`(1,64,128,128)` | 0.00e+00 | OK |
   | rank-0 scalar slope | `(2,3,4,5)` | `(2,3,4,5)` | 0.00e+00 | OK |
   | same-rank multi non-broadcast, `slope == X` | `(2,3,4,5)` | `(2,3,4,5)` | 
0.00e+00 | OK |
   | same-rank multi non-broadcast `(1,3,4,5)` | `(2,3,4,5)` | `(2,3,4,5)` | 
0.00e+00 | OK |
   | lower-rank multi non-broadcast `(4,5)` | `(2,3,4,5)` | `(2,3,4,5)` | 
0.00e+00 | OK |
   | same-rank single non-broadcast `(1,1,4,1)` (regression) | `(2,3,4,5)` | 
`(2,3,4,5)` | 0.00e+00 | OK |
   | 1-D `(5,)` / all-ones (regression) | `(2,3,4,5)` | `(2,3,4,5)` | 0.00e+00 
| OK |
   
   Total: **113 valid cases (onnxruntime-accepted), 113 OK, 0 rejected, 0 
numeric
   mismatches** (pre-fix: 46 OK / 67 rejected). The motivating Real-ESRGAN case
   `X(1,64,128,128)` + `slope(64,1,1)` now imports and matches onnxruntime 
exactly.
   Run:
   
   ```bash
   python results/TVM/deepseek-v4-flash/prove_hum/onnx_PRelu/5修复_差分验证.py
   ```
   
   ## Files changed
   
   - `python/tvm/relax/frontend/onnx/onnx_frontend.py` — `PRelu._impl_v1`: 
elementwise
     `where(x < 0, s * x, x)` fallback for multi-non-broadcast slopes.
   


-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to