Files
apache--tvm/python/tvm/ir/_overload_prim_expr.py
Tianqi Chen 275114b327 [REFACTOR][IR] Unify PrimExpr with Expr typed view (#19910)
## Summary
- Make `PrimExpr` a typed C++ view over `Expr` values whose
`ExprNode::ty` is `PrimType`, instead of using a separate runtime node
class as the proof of primitive-ness.
- Use the shared `ir::Call` node for Relax, TIRX, and primitive-valued
calls, while keeping primitive-only APIs explicit at their semantic
boundaries.
- Keep Python on the general `Expr` surface for primitive-typed values
so `isinstance` behavior does not imply a nominal primitive-expression
subclass.

## Design Rationale
The main advantage of this change is that common expression nodes such
as `Call` can be unified without specializing each one to `PrimType`. A
single `ir::Call` can represent a Relax tensor call, a Relax scalar
call, or a primitive-valued intrinsic call; the result type stored in
`ExprNode::ty` determines whether that particular value can be viewed as
`PrimExpr`.

This keeps the IR node hierarchy focused on expression structure rather
than result-type categories. Nodes that are intrinsically primitive,
such as integer and floating-point literals or TIRX primitive operators,
still have strongly typed C++ APIs and data structures. General nodes
whose result type may vary, such as `Call`, remain general `Expr` nodes
and are narrowed to `PrimExpr` only where primitive-only semantics are
required.

The PR also keeps the compatibility surface practical: C++
primitive-only APIs continue to accept `PrimExpr`, Python exposes a
compatibility predicate for checking the primitive typed category, and
visitors/printers use one natural `Call` path rather than duplicating
Relax and primitive call handling. Missing expression types are
represented explicitly with `Type::Missing()` so constructors can leave
type inference to later analysis without relying on nullable `Type`
values.
2026-07-01 18:55:33 -04:00

154 lines
2.6 KiB
Python

# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
"""Primitive-expression overloads for shared IR expressions."""
def __add__(_lhs, _rhs):
return NotImplemented
def __radd__(_lhs, _rhs):
return NotImplemented
def __sub__(_lhs, _rhs):
return NotImplemented
def __rsub__(_lhs, _rhs):
return NotImplemented
def __mul__(_lhs, _rhs):
return NotImplemented
def __rmul__(_lhs, _rhs):
return NotImplemented
def __div__(_lhs, _rhs):
return NotImplemented
def __rdiv__(_lhs, _rhs):
return NotImplemented
def __truediv__(_lhs, _rhs):
return NotImplemented
def __rtruediv__(_lhs, _rhs):
return NotImplemented
def __floordiv__(_lhs, _rhs):
return NotImplemented
def __rfloordiv__(_lhs, _rhs):
return NotImplemented
def __mod__(_lhs, _rhs):
return NotImplemented
def __rmod__(_lhs, _rhs):
return NotImplemented
def __neg__(_value):
return NotImplemented
def __lshift__(_lhs, _rhs):
return NotImplemented
def __rlshift__(_lhs, _rhs):
return NotImplemented
def __rshift__(_lhs, _rhs):
return NotImplemented
def __rrshift__(_lhs, _rhs):
return NotImplemented
def __and__(_lhs, _rhs):
return NotImplemented
def __rand__(_lhs, _rhs):
return NotImplemented
def __or__(_lhs, _rhs):
return NotImplemented
def __ror__(_lhs, _rhs):
return NotImplemented
def __xor__(_lhs, _rhs):
return NotImplemented
def __rxor__(_lhs, _rhs):
return NotImplemented
def __invert__(_value):
return NotImplemented
def __lt__(_lhs, _rhs):
return NotImplemented
def __le__(_lhs, _rhs):
return NotImplemented
def __eq__(_lhs, _rhs):
return NotImplemented
def __ne__(_lhs, _rhs):
return NotImplemented
def __gt__(_lhs, _rhs):
return NotImplemented
def __ge__(_lhs, _rhs):
return NotImplemented
def equal(_lhs, _rhs, _span=None):
return NotImplemented
def astype(_value, _dtype, _span=None):
return NotImplemented