Skip to content

[Feature] Remove scalar.cast_index and infer index-to-PTO-scalar conversions in PTODSL #1081

Description

@Zhendong404

Summary

希望移除 PTODSL 的 scalar.cast_index 公共接口,由 DSL 在需要 PTO scalar 的上下文中自动完成 index 到对应 PTO scalar 类型的转换。

当前前端生成或用户编写 PTODSL 时,需要显式插入类似下面的转换:

value = scalar.cast_index(index_value, dtype)

这将 MLIR/编译器内部的 index 类型适配细节暴露给了 DSL 用户和上游代码生成器。对用户而言,这种转换通常没有业务语义,只是为了满足 PTO op 的 scalar operand 类型要求。

Motivation / use case

  • 简化 PTODSL API,避免用户理解并手动处理 MLIR index 与 PTO scalar 类型之间的边界。
  • 减少上游前端(例如 TileLang)为每个 scalar operand 特判并生成 scalar.cast_index 的代码。
  • 让索引表达式可以自然地用于 offset、shape、stride、loop-derived scalar 等参数,同时由 DSL 根据目标 operand 的期望类型完成适配。
  • 将转换规则集中在 PTODSL 类型系统/operation builder 中,避免不同调用方使用不一致的 cast 逻辑。

Proposed API / behavior

当一个 index 值传给期望 PTO scalar 类型的 operand 时,PTODSL 应自动插入合法转换。例如:

# Before
pto.some_op(scalar.cast_index(i, pto.i32))

# After
pto.some_op(i)

建议的行为边界:

  1. 根据 operation schema 中 operand 的期望类型决定目标 scalar 类型,而不是由用户重复指定 dtype。
  2. 仅在目标类型明确且转换合法时自动转换;目标类型不明确或存在歧义时给出清晰诊断。
  3. 对窄化转换明确规则:如果 index 到目标整数类型可能丢失范围信息,应采用项目统一的语义,或要求能够证明安全,否则报错。
  4. 保持转换在生成的 PTO/MLIR IR 中显式可见,方便验证和调试;这里只移除用户侧显式 API 调用,并非隐藏 IR cast。
  5. 在迁移期可保留 scalar.cast_index 并标记 deprecated,待上游调用方迁移后再删除。

Alternatives considered

继续要求上游前端或 PTODSL 用户显式调用 scalar.cast_index。这种方式实现直接,但会把底层类型适配细节扩散到每个调用方,增加生成代码噪音和维护成本,也容易因目标 dtype 选择不一致引入问题。

Additional context

建议增加以下测试:

  • index 作为不同 PTO op 的 scalar operand 时能自动转换。
  • operation schema 能唯一确定 i32/i64 等目标类型。
  • 类型不明确、非法转换和可能不安全的窄化转换产生明确诊断。
  • 已经是正确 PTO scalar 类型的 operand 不产生冗余 cast。
  • scalar.cast_index 的兼容/弃用路径有覆盖。

验收标准:上游生成的 PTODSL 源码无需出现 scalar.cast_index,仍能生成类型合法、转换语义明确的 PTO IR。

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions