Slang 自动微分类型系统深度解析:IDifferentiable、DifferentialPair 与可微类型检查机制

发布时间:2026/9/17 14:27:41
Slang 自动微分类型系统深度解析:IDifferentiable、DifferentialPair 与可微类型检查机制 Slang 自动微分类型系统深度解析IDifferentiable、DifferentialPair 与可微类型检查机制【免费下载链接】slangMaking it easier to work with shaders项目地址: https://gitcode.com/GitHub_Trending/sl/slang本文面向 Slang 编译器贡献者从编译器工程视角剖析自动微分autodiff类型系统的核心组件IDifferentiable接口、DifferentialPairT类型、可微类型的自动合成、以及fwd_diff/bwd_diff高阶调用的类型检查流程。读完本文你将理解 Slang 如何在类型层面支撑一阶与高阶自动微分以及no_diff、detach()等机制如何控制导数传播。普通 Slang 用户可先参阅 用户指南动手深入之前建议先阅读自动微分基础的 Basics 文档掌握 primal / differential / 前向与反向模式的基本概念。类型系统的四大组成部分自动微分在 Slang 中并不是一个游离在类型系统之外的魔法流程而是建立在精心设计的类型系统之上。docs/design/autodiff/types.md将这一类型系统拆解为四个主要部分IDifferentiable接口定义“什么是可微类型”是核心模块与用户代码共同依赖的基石DifferentialPairT类型用一个类型同时携带原值primal与对应的微分值differential自动微分算子的类型检查fwd_diff/bwd_diff本质上是高阶函数因此前端的语义检查需要一套专门的高阶函数检查机制导数数据流分析通过静态分析警告用户“意外中断导数传播”的情况并借助no_diff装饰明确地切断导数。下面逐一对这四个部分展开并结合仓库源码印证其实现细节。interface IDifferentiable可微类型的契约IDifferentiable定义在核心模块 source/slang/core.meta.slang 中第 648 行起是标记“可微类型”的基础接口既供核心模块内部使用也向用户代码开放。其定义刻意只封装以下 4 个要素Differential该类型微分值的类型。允许用户自定义数据结构来承载微分值——相比完全依赖编译器合成用户可以针对空间占用做优化。dadd(Differential, Differential) - Differential两个微分值相加。由于导数计算本质上是线性的我们只需这一种二元运算。其实现必须满足结合律与交换律否则生成的导数代码可能是错误的。dzero() - Differential加法单位元零值用于在梯度聚合时初始化累加变量。dmulS: __BuiltinRealType(S, Differential)实数与微分类型的标量乘法实现必须对微分加法dadd满足分配律。第 2、3、4 点共同来源于向量空间的概念——任何 Slang 函数的导数值总是构成一个向量空间因此只需要线性运算即可表达全部导数语义。值得注意的细节是仓库中的实际接口定义source/slang/core.meta.slang第 648-663 行将dmul视作编译器内部合成的能力接口显式声明的成员只有三个且都用__builtin_requirement(...)标记了内置需求编号__magic_type(DifferentiableType) [KnownBuiltin($( (int)KnownBuiltinDeclName::IDifferentiable))] interface IDifferentiable { // 注意编译器实现要求 Differential 关联类型必须最先定义。 __builtin_requirement($((int)BuiltinRequirementKind::DifferentialType)) associatedtype Differential : IDifferentiable; /// 返回微分类型的零初始化值。 __builtin_requirement($((int)BuiltinRequirementKind::DZeroFunc)) [Differentiable] static Differential dzero(); /// 将两个微分值相加并返回结果。 __builtin_requirement($((int)BuiltinRequirementKind::DAddFunc)) [Differentiable] static Differential dadd(Differential, Differential); };此外接口还隐含一条二阶约束T.Differential可以不同于T但T.Differential.Differential必须等于T.Differential自身。这条规则保证了二阶及更高阶导数仍然落在同一微分类型上是支撑高阶自动微分的关键。微分成员关联[DerivativeMember]装饰器在有些场景下编译器需要知道原始类型的字段如何映射到微分类型的字段。典型场景是通过花括号{}隐式构造结构体IR 中对应kIROp_MakeStruct时的微分处理。为此 Slang 提供了[DerivativeMember(DifferentialTypeName.fieldName)]装饰器来显式标记这种关联关系其 AST 节点为DerivativeMemberAttribute见 source/slang/slang-ast-modifier.h 第 1839 行IR 表示为kIROp_DerivativeMemberDecoration。示例struct MyType : IDifferentiable { typealias Differential MyDiffType; float a; [DerivativeMember(MyDiffType.db)] float b; /* ... */ }; struct MyDiffType { float db; };其中b的导数会存放在MyDiffType.db中而a未标注如果它的类型本身可微编译器会为其自动选择微分字段。编译器在语义检查阶段会校验该属性的合法性只有可微类型结构的成员才能使用DerivativeMemberAttribute非法使用会产生诊断错误参见checkDerivativeMemberAttributeParent位于 source/slang/slang-check-decl.cpp。聚合类型的IDifferentiable一致性自动合成要求用户为每个自定义struct手写关联的Differential类型、字段映射和三个接口方法非常繁琐。对于聚合类型struct / tuple / 数组等这些实现可以通过分析其成员类型是否满足IDifferentiable来自动构造。合成流程大致如下IDifferentiable的各需求组件被标记上特殊的__builtin_requirement(unique_integer_id)装饰器携带BuiltinRequirementKind枚举值。当检查类型与接口的一致性conformance时如果用户提供的定义无法满足某个带内置标记的需求编译器会分派到trySynthesizeRequirementWitness执行合成。对于用户自定义类型Differential类型在一致性检查期间通过trySynthesizeDifferentialAssociatedTypeRequirementWitnesssource/slang/slang-check-decl.cpp 第 3804 行合成逐字段检查每个成员类型是否满足IDifferentiable查找其对应的Differential类型再用这些微分类型构造新的聚合类型。由于某个成员类型的Differential可能尚未合成查找系统trySynthesizeRequirementWitness会先合成一个带ToBeSynthesizedModifier定义于 source/slang/slang-ast-modifier.h 第 173 行的临时空类型待成员类型完成一致性检查后再回填字段。对于用户自定义类型dadd、dzero和dmul方法在trySynthesizeDifferentialMethodRequirementWitnesssource/slang/slang-check-decl.cpp 第 9578 行中合成利用Differential成员及其[DerivativeMember]装饰器确定需要考虑哪些字段、每个字段使用哪个基础类型。合成有两种模式完全归纳模式fully-inductive用于dadd和dzero即对Differential类型的各字段分别调用dadd/dzero并组装结果// 由 struct T {FT1 field1; FT2 field2;} 合成 T.Differential dadd(T.Differential a, T.Differential b) { return Differential( FT1.dadd(a.field1, b.field1), FT2.dadd(a.field2, b.field2), ) }固定首参模式fixed-first arg用于dmul。因为第一个参数是公共标量只需对其余参数做归纳// 由 struct T {FT1 field1; FT2 field2;} 合成 T.Differential dmulS:__BuiltinRealType(S s, T.Differential a) { return Differential( FT1S.dmul(s, a.field1), FT2S.dmul(s, a.field2), ) }在自动微分过程中编译器有时会合成新的聚合类型最常见的是中间上下文类型kIROp_BackwardDerivativeIntermediateContextType它在自动微分 pass 完成后被降级为普通 struct。由于这些类型可能被进一步微分高阶自动微分必须为它们合成IDifferentiable一致性。这部分实现在fillDifferentialTypeImplementationForStruct(...)中逻辑与 AST 侧的合成大致类似。此外source/slang/core.meta.slang 第 520-547 行的文档还补充了一条重要的自动满足优化规则如果合成出的Differential类型与原始类型字段完全相同且各字段类型也一致那么编译器会直接用原始类型自身充当Differential类型即T.Differential T而不再创建新类型。标量类型float、向量float3等正是通过这一规则实现“自身即微分类型”的。可微类型字典Differentiable Type Dictionaries自动微分过程中IR 各 pass 需要频繁查询“某个IRType是否可微”并获取对应IDifferentiable方法的引用。这些查询还必须能作用于泛型参数定义在泛型容器内部和以接口类型作为参数的 existential 类型。为了覆盖这些不同的类型系统Slang 采用了一套“类型字典”机制为每个函数关联一份相关类型的字典在对标记为[Differentiable]的函数内表达式调用CheckTerm()时检查解析出的类型是否满足IDifferentiable。若满足就把该类型连同其可微性 witness 加入字典字典目前存放在与该[Differentiable]修饰符对应的DifferentiableAttribute上。降级到 IR 时创建DifferentiableTypeDictionaryDecoration持有字典中所有类型的 IR 版本以及它们IDifferentiablewitness 表的引用。合成导数代码时所有 transcriber pass 通过DifferentiableTypeConformanceContext::setFunc()加载类型字典。随后DifferentiableTypeConformanceContext提供便捷函数用于查询可微类型、获取合适的IDifferentiable方法、构造合适的DifferentialPairT。泛型类型上的微分信息查找泛型定义的类型同样会被放入可微类型字典但它们的 witness 表本身是参数而不是具体的 witness 表。当自动微分 pass 需要查找微分类型或调用IDifferentiable方法时会被转成对 witness 表参数的查找即Lookup(InterfaceRequirementKey, WitnessTableParameter)。注意这些查找指令被插入到泛型父容器中而不是最内层的函数里。示例T myFuncT:IDifferentiable(T a) { return a * a; } // 反向模式微分版本 void bwd_myFuncT:IDifferentiable( inout DifferentialPairT dpa, T.Differential dOut) // T.Differential 即 Lookup(Differential, T_Witness_Table) { T.Differential da T.dzero(); // T.dzero 即 Lookup(dzero, T_Witness_Table) da T.dadd(dpa.p * dOut, da); // T.dadd 即 Lookup(dadd, T_Witness_Table) da T.dadd(dpa.p * dOut, da); dpa diffPair(dpa.p, da); }existential 类型上的微分信息查找existential 类型是“以接口为类型”的值运行时可能存在多种实现。existential 值在运行时携带具体类型信息本质上是一种“带标签的联合类型”tagged union。existential 的微分类型existential 的微分类型定义起来比较棘手因为类型系统对.Differential的唯一约束是“它也满足IDifferentiable”。因此任何满足IInterface : IDifferentiable的接口其微分类型都是接口IDifferentiable本身。这带来一个问题Slang 通常要求一个静态的anyValueSize它必须是所有满足类型尺寸的严格上界用于为联合类型分配空间。由于IDifferentiable定义在核心模块core.meta.slang中且用户也可使用无法可靠地定义一个静态上界。为此 Slang 新增了一个any-value-size 推断 passslang-ir-any-value-inference.h/slang-ir-any-value-inference.cpp位于 source/slang 目录它在最终链接后的 IR 中收集“满足每个接口的类型清单”从而确定一个相关的上界。这样做可以忽略那些满足IDifferentiable但未在最终 IR 中使用的类型得到更紧凑的上界。未来工作这一方案虽然可用但存在局部性问题IDifferentiable的尺寸是可见模块中所有满足IDifferentiable类型尺寸的最大值而实际上我们只关心那些作为T : IInterface的T.Differential出现的类型子集。原因在于执行关联类型查找后Slang IR 丢弃了查找起点的基础接口信息只考虑约束接口这里即Differential : IDifferentiable。解决思路包括(i) 静态分析每个使用位置的可能类型集合并传播以收窄类型范围或 (ii) 引入泛型参数化接口例如IDifferentiableT使每个版本拥有不同的满足类型集合。示例伪代码对应 IR部分指令还会进一步降级interface IInterface : IDifferentiable { [Differentiable] This foo(float val); [Differentiable] float bar(); }; float myFunc(IInterface obj, float a) { IInterface k obj.foo(a); return k.bar(); } // 反向模式微分版本伪代码 void bwd_myFunc( inout DifferentialPairIInterface dpobj, inout DifferentialPairfloat dpa, float.Differential dOut) // T.Differential 即 Lookup(Differential, T_Witness_Table) { // 前向primalpass.. IInterface obj dpobj.p; IInterface k obj.foo(a); // ..... // 反向backwardpass DifferentialPairIInterface dpk diffPair(k); bwd_bar(dpk, dOut); IDifferentiable dk dpk.d; // IInterface 的微分类型即 IDifferentiable DifferentialPairIInterface dp diffPair(dpobj.p); bwd_foo(dpobj, dpa, dk); }existential 上的dadd()与dzero()查找对 existential 类型的查找分两种情况。更常见的是封闭盒closed-boxexistential即直接以接口表示这种类型的每个值都携带类型标识符、witness 表标识符以及值本身。较少见的情况是函数调用直接作用于被转换cast成具体类型后的值上。封闭 existential 的dzero()NullDifferential类型对于具体类型乃至泛型类型我们可以调用对应的Type.dzero()来初始化导数累加变量。但对 existential 微分当前类型为IDifferentiable却不行——我们还必须把 existential 的类型 id 初始化为某个具体实现而运行前我们并不知道是哪一个这是一个只有在第一个微分值产生后才可知的运行时值。为此Slang 声明了一个特殊类型NullDifferential充当任何IDifferentiableexistential 对象的“none 类型”。其定义位于 source/slang/diff.meta.slang 第 178-193 行// 一个充当“零微分”运行时哨兵值的 none-type主要用于内部使用。 [__AutoDiffBuiltin] [KnownBuiltin($((int)KnownBuiltinDeclName::NullDifferential))] struct NullDifferential : IDifferentiable { // 暂时至少保留一个字段确保类型非空 float dummy; typedef NullDifferential Differential; [Differentiable] [ForceInline] static Differential dzero() { return { 0.0f }; } [Differentiable] [ForceInline] static Differential dadd(Differential, Differential) { return { 0.0f }; } };封闭 existential 的dadd()__existential_dadd我们不能直接对两个IDifferentiable类型的 existential 微分调用dadd()因为必须处理“其中一个操作数是NullDifferential”的情况而dadd()只对同类型的微分有定义。Slang 目前的处理方式是合成一个特殊方法__existential_dadd即getOrCreateExistentialDAddMethod位于 source/slang/slang-ir-autodiff.cpp 第 403 行。该方法的实现逻辑如下与源码中的 IR 构造一一对应提取第一个操作数 a 的 existential 类型与 witness 表通过emitIsType检查其是否为NullDifferential若是直接返回 b否则检查第二个操作数 b 是否为NullDifferential若是直接返回 a若两者都非空则从 a 的 witness 表上按dadd需求键查找具体类型的dadd方法提取两边的实际值并调用最后用emitMakeExistential将结果重新包装成IDifferentiable类型的 existential 返回。也就是说__existential_dadd在运行时做类型 id 检查若任一操作数是NullDifferential则返回另一个若都不是则分派到具体类型的dadd。这是对“封闭盒”语义的运行时动态分派。开放openexistential 的dadd()与dzero()如果操作的是具体类型的值即通过ExtractExistentialValue(ExistentialParam)打开的 existential 值那么可以像泛型一样进行查找。所有 existential 参数都携带 witness 表编译器插入提取 witness 表的指令并据此查找即dadd使用Lookup(dadd, ExtractExistentialWitnessTable(ExistentialParam))并对查找结果发起调用。struct DifferentialPairT: IDifferentiable原值与微分的“成对”载体第二个核心组件是DifferentialPairT:IDifferentiable表示“一个原值 其对应微分值”的配对。其用途主要有二在合成出的导数方法之间传递/接收导数以及作为 IR 侧的块参数block parameter。由于fwd_diff(fn)与bwd_diff(fn)都是“函数到函数”的变换Slang 前端会把fn的类型翻译成其导数版本以便对调用参数做类型检查。DifferentialPair在 source/slang/core.meta.slang 第 778-794 行定义如下__genericT : IDifferentiable __magic_type(DifferentialPairType) __intrinsic_type($(kIROp_DifferentialPairType)) struct DifferentialPair : IDifferentiable { typedef DifferentialPairT.Differential Differential; typedef T.Differential DifferentialElementType; __intrinsic_op($(kIROp_MakeDifferentialPair)) [Differentiable] __init(T _primal, T.Differential _differential); property p : T { __intrinsic_op($(kIROp_DifferentialPairGetPrimal)) get; } // ... property d : T.Differential 对应 kIROp_DifferentialPairGetDifferential };即DifferentialPairT自身的Differential是DifferentialPairT.Differential再次印证“二阶微分类型不变”的规则并且它自己也满足IDifferentiable——这正是高阶自动微分能够对“携带微分的值”继续求导的基础。成对类型的降级Pair Type LoweringDifferentialPair在 AST 与 IR 各 pass 中都是特殊类型AST 节点DifferentialPairTypeIR 为kIROp_DifferentialPairType因为它被前端语义检查和导数代码合成反复使用。一旦自动微分 pass 全部完成成对类型会被降级成简单struct以便各后端正常发射这项工作由DiffPairLoweringPass完成位于 source/slang/slang-ir-autodiff-pairs.cpp 第 458 行。与此配套还定义了成对构造与提取指令kIROp_MakeDifferentialPair构造、kIROp_DifferentialPairGetDifferential与kIROp_DifferentialPairGetPrimal提取它们分别被降级为 struct 构造与字段访问。“用户代码”成对类型User-code Differential Pairs既然成对类型因为 IR pass 中的特殊处理而使用专门的 IR 指令那么反过来有些场景希望自动微分 pass把成对类型当作普通 struct 类型处理。这主要发生在高阶自动微分中——用户希望对同一段代码多次求导。Slang 的做法是在每一轮自动微分迭代结束时把所有相关的成对类型改写为“无关irrelevant成对类型”kIROp_DifferentialPairUserCode以及“无关访问器”kIROp_DifferentialPairGetDifferentialUserCode、kIROp_DifferentialPairGetPrimalUserCode这样下一轮迭代就会把它们当作普通可微类型。这些用户代码版本同样会被降级为 struct。自动微分调用的类型检查及其它高阶函数fwd_diff与bwd_diff被表示成“输入一个函数、返回其导数函数”的高阶函数因此前端语义检查需要某种高阶函数概念才能检查和降级这类调用。高阶调用基类HigherOrderInvokeExpr所有高阶变换都派生自HigherOrderInvokeExprsource/slang/slang-ast-expr.h 第 803 行。自动微分有两种表达式类ForwardDifferentiateExpr第 826 行与BackwardDifferentiateExpr第 835 行二者都派生自该父类表达式。高阶函数调用检查HigherOrderInvokeExprCheckingActions在 Slang 中解析具体的方法并非易事——它支持重载、类型强制转换等特性而当函数变换出现在调用链中时问题更加复杂。例如对于fwd_diff(f)(DiffPairfloat(...), DiffPairdouble(...))我们需要根据变换后的参数类型找到f的正确匹配。为此使用如下工作流实现在 source/slang/slang-check-expr.cpp 的HigherOrderInvokeExprCheckingActions第 5817 行起HigherOrderInvokeExprCheckingActions基类为不同的高阶表达式提供实现其类型翻译的机制即“变换后的函数是什么类型”。前向/反向分别由ForwardDifferentiateExprCheckingActions第 5888 行与BackwardDifferentiateExprCheckingActions第 5929 行实现。检查机制把所有检测到的f的重载候选逐一通过类型翻译用结果组装出一个新候选组这些新函数是“临时的”。这个新候选组被ResolveInvoke用来结合用户提供的实参列表做重载决议与类型强制转换。解析出的签名若有随后被替换为对应的函数引用并包装进相应的高阶 invoke 中。示例假设有两个同名函数f签名分别为int - float和double, double - float我们需要解析调用fwd_diff(f)(DiffPairfloat(1.0, 0.0), DiffPairfloat(0.0, 1.0))。高阶检查动作会合成“临时”翻译签名组int - DiffPairfloat与DiffPairdouble, DiffPairdouble - DiffPairfloat。Invoke 解析随后通过把float自动转换为double把候选收窄到唯一匹配DiffPairdouble, DiffPairdouble - DiffPairfloat。解析完成后返回InvokeExpr(ForwardDifferentiateExpr(f : double, double - float), casted_args)——即把对应函数包装进对应的高阶表达式。属性化类型no_diff参数出于正确性考虑经常需要阻止梯度穿过某些参数。例如随机样本的值通常不应被微分否则数学结果可能不正确。即使参数类型满足IDifferentiableSlang 也提供no_diff操作符将参数标记为不可微float myFunc(float a, no_diff float b) { return a * b; } // 得到的前向模式导数 DiffPairfloat myFunc(DiffPairfloat dpa, float b) { return diffPair(dpa.p * b, dpa.d * b); }可以看到b在导数版本中保持普通float不参与微分。Slang 在 IR 侧使用OpAttributedType表示这类参数的类型上例中b降级后的类型是OpAttributedType(OpFloat, OpNoDiffAttr)在前端则用ModifiedTypeAST 节点表示定义于 source/slang/slang-ast-type.h 第 1364 行。有时这层附加信息会干扰类型相等性检查等“与no_diff无关”的机制因此 Slang 提供了unwrapAttributedType辅助函数source/slang/slang-ir-util.h 第 314 行来剥离属性化类型层。导数数据流分析Slang 还有一个导数数据流分析 pass它在函数降级到 IR 之后、链接步骤之前对每个函数单独执行实现于 source/slang/slang-ir-check-differentiability.cpp 与同名头文件。该 pass 的职责是强制保证可微类型的指令会传播导数除非用户通过detach()或no_diff显式丢弃导数。原因在于Slang 要求函数必须标注[Differentiable]才允许传播导数否则该函数被视为不可微其导数实际为 0。这会导致令人沮丧的场景——函数可能在无意中丢弃导数。例如float nonDiffFunc(float x) { /* ... */ } float differentiableFunc(float x) // 忘记标注 [Differentiable] { /* ... */ } float main(float x) { // 用户没有意识到“本应可微的函数”并未被微分 // 因为这里的类型全是 float。 // return nonDiffFunc(x) * differentiableFunc(x); }数据流分析强制要求在可微上下文中使用的不可微函数必须显式丢弃其导数。这样用户就能清楚地知道某个调用是被微分了还是被丢弃了。同一个例子加上no_diff强制约束后float nonDiffFunc(float x) { /* ... */ } [Differentiable] float differentiableFunc(float x) { /* ... */ } float main(float x) { return no_diff(nonDiffFunc(x)) * differentiableFunc(x); }no_diff只能直接用于函数调用上它会变成一个TreatAsDifferentiableDecoration表示该函数不会产生导数。导数数据流分析的工作方式与标准数据流分析类似先组装一个“产生导数”的指令集合从可微类型且无显式no_diff的参数出发沿块内每条指令传播——只要某条指令的操作数携带导数、且结果类型可微该指令就携带导数。再组装一个“期望导数”的指令集合这些是可微函数中未被no_diff标记的可微操作数。然后对该集合做反向传播不断加入所有可微操作数并重复此过程。在反向传播过程中如果“期望”集合里存在某个OpCall不在“产生”集合中就意味着梯度未被显式丢弃此时为用户生成一条诊断信息。正是这最后一步把“隐式丢弃导数”变成了“显式、可诊断”的行为从类型系统层面保证了自动微分代码的可预期性。小结Slang 的自动微分类型系统可以概括为一条清晰的链条IDifferentiable定义可微契约Differential/dadd/dzero/dmul编译器为聚合类型自动合成一致性实现DifferentialPairT统一携带原值与微分值并作为高阶导数的基础载体类型字典机制把“可微性”从具体类型扩展到泛型与 existential 类型HigherOrderInvokeExprCheckingActions让fwd_diff/bwd_diff作为高阶函数可以正确完成重载解析与类型转换no_diff与导数数据流分析则共同保证导数传播的显式性与正确性。这些机制的实现横跨 source/slang/core.meta.slang接口与内置类型定义、source/slang/diff.meta.slangNullDifferential、diffPair等辅助定义、source/slang/slang-check-decl.cpp一致性合成、source/slang/slang-check-expr.cpp高阶调用检查、source/slang/slang-ir-autodiff.cpp__existential_dadd等运行时支持、source/slang/slang-ir-autodiff-pairs.cpp成对类型降级以及 source/slang/slang-ir-check-differentiability.cpp数据流分析。对自动微分在用户层面的使用方式感兴趣可继续阅读 自动微分用户指南。【免费下载链接】slangMaking it easier to work with shaders项目地址: https://gitcode.com/GitHub_Trending/sl/slang创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考