-
Notifications
You must be signed in to change notification settings - Fork 101
FlagTree Backend Specialization
FlagTree 设计的后端统一特化,目的是整合后端接入范式,清晰化管理后端的特化实现,为后端维护特有变更、复用 FlagTree 公共能力、升级 Triton 版本、避免编辑模式出错等提供了工程基础。
具体实施方案是将各后端对 Triton 的特化,从以往的直接修改主干代码或者是将整套代码放在后端目录,标准化为在复用主干代码的基础上,在后端目录中给出差异化实现。
主干代码在原则上,既不允许直接出现某后端的特化实现,也不允许对后端做选择判断后特化实现,除非该主干代码文件原先已有后端选择实现。
python/triton/ 目录下的主干代码,可通过调用接口,连接到 third_party/${backend_name}/backend/spec/triton/ 目录下的特化实现。接口定义在 python/triton/flagtree_spec.py。
spec_path 接口原理:在 __init__.py 劫持 __path__,优先在后端特化路径查找同名文件(模块),若不存在才使用主干代码
案例:https://github.com/flagos-ai/FlagTree/pull/687
目标:mthreads 后端需要在 python/triton/tools/ 另行实现 compile.py、link.py、ragged_tma.py,目录下其他文件复用主干代码
特化代码:在 third_party/mthreads/backend/spec/triton/runtime/ 添加与主干代码不同的 compile.py、link.py、ragged_tma.py,不添加与主干代码完全相同或功能完全被主干代码包含的文件,也不必添加 __init__.py
主干目录初始化代码:python/triton/tools/__init__.py 调用接口
from triton.flagtree_spec import spec_path
# flagtree backend path specialization
spec_path(__path__)
...spec 接口原理:通过 third_parth/{backend_name}/backend/spec/triton/__init__.py 调用后端定义的函数,该功能常用于在 __init__.py 修改模块的导出符号,也可用于其他 py 文件
案例:https://github.com/flagos-ai/FlagTree/pull/687
目标:mthreads 后端需要在 python/triton/language/__init__.py 添加导出 squeeze, unsqueeze 等符号
特化代码:新增文件 third_party/mthreads/backend/spec/triton/language/__init__.py 实现 language_extend_globals 函数,注意其中的 import 必须使用绝对路径
def language_extend_globals(globals_dict):
# NOTE: Must use absolute path import.
from triton.language.standard import squeeze, unsqueeze
globals_dict["squeeze"] = squeeze
globals_dict["unsqueeze"] = unsqueeze特化顶层初始化代码:third_party/mthreads/backend/spec/triton/__init__.py 导入该函数
from .language import language_extend_globals主干目录初始化代码:python/triton/language/__init__.py 通过接口调用 language_extend_globals 函数,注意用于修改模块的导出符号时,spec 应放在所有 import 之后
from triton.flagtree_spec import spec
... # import
# flagtree backend specialization
spec("language_extend_globals", globals())
__all__ = [ ... ]spec_func 接口原理:通过 third_parth/{backend_name}/backend/spec/triton/__init__.py 返回后端新增或修改的函数定义,可用于任意的 py 文件
案例:https://github.com/flagos-ai/FlagTree/pull/687
目标:mthreads 后端需要在 python/triton/_utils.py 添加 apply_with_path 的函数定义
注意:由于 python/triton/ 是所有目录的父目录,不允许使用 spec_path 劫持 __path__,因此无法整文件特化 _utils.py,转而使用 spec_func 添加函数定义
特化顶层初始化代码:third_party/mthreads/backend/spec/triton/__init__.py 导入该函数
from ._utils import apply_with_path主干目录代码:python/triton/_utils.py 通过接口导入 apply_with_path 的函数定义
from triton.flagtree_spec import spec_func
# flagtree backend func specialization
apply_with_path = spec_func("apply_with_path")