275114b327
## 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.
83 lines
2.5 KiB
Python
83 lines
2.5 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.
|
|
"""Target dependent intrinsic registration."""
|
|
|
|
from tvm.ir import register_intrin_lowering
|
|
from tvm.tirx import call_pure_extern
|
|
|
|
|
|
def _rule_float_suffix(op):
|
|
"""Intrinsic rule: Add float suffix if it is float32.
|
|
|
|
This is an example intrinsic generation rule.
|
|
|
|
Parameters
|
|
----------
|
|
op : Expr
|
|
The call expression of original intrinsic.
|
|
|
|
Returns
|
|
-------
|
|
ret : Expr
|
|
The translated intrinsic rule.
|
|
Return same op if no translation is possible.
|
|
|
|
See Also
|
|
--------
|
|
register_intrin_lowering : The registration function for intrinsic lowering rule.
|
|
"""
|
|
name = op.op.name
|
|
assert name.startswith("tirx.")
|
|
prefix = name[4:]
|
|
|
|
if op.ty.dtype == "float32":
|
|
return call_pure_extern(op.ty, f"{prefix}f", *op.args)
|
|
if op.ty.dtype == "float64":
|
|
return call_pure_extern(op.ty, prefix, *op.args)
|
|
return op
|
|
|
|
|
|
def _rule_float_direct(op):
|
|
"""Intrinsic rule: Directly call pure extern function for floats.
|
|
|
|
This is an example intrinsic generation rule.
|
|
|
|
Parameters
|
|
----------
|
|
op : Expr
|
|
The call expression of original intrinsic.
|
|
|
|
Returns
|
|
-------
|
|
ret : Expr
|
|
The translated intrinsic rule.
|
|
Return same op if no translation is possible.
|
|
|
|
See Also
|
|
--------
|
|
register_intrin_lowering : The registration function for intrinsic lowering rule.
|
|
"""
|
|
if str(op.ty.dtype).startswith("float"):
|
|
return call_pure_extern(op.ty, op.op.name[4:], *op.args)
|
|
return None
|
|
|
|
|
|
# opencl pattern for exp
|
|
register_intrin_lowering("tirx.exp", target="opencl", f=_rule_float_direct, level=99)
|
|
# default pattern for exp
|
|
register_intrin_lowering("tirx.exp", target="default", f=_rule_float_suffix, level=99)
|