From 131352da024990570479be579981008c970baceb Mon Sep 17 00:00:00 2001 From: Hiroshi Shinaoka Date: Sat, 11 Apr 2026 21:02:04 +0900 Subject: [PATCH 1/2] feat: add ADContext associated type to PrimitiveOp trait Add `type ADContext: Default` to PrimitiveOp. Both `linearize` and `transpose_rule` now receive `&mut Self::ADContext`, enabling AD rules to access runtime context (e.g., concrete tensor shapes, guard recording) during graph-to-graph differentiation. Existing impls use `type ADContext = ()` for zero-cost backward compat. New tests verify: - RecordingOp: context records linearize/transpose calls - BranchOp: context controls which ops are emitted (simulating SVD's `if m > n` branching pattern) - Multiple branch points: guards accumulate across sequential ops Co-Authored-By: Claude Opus 4.6 (1M context) --- src/primitive_op.rs | 12 + tests/trait_tests.rs | 518 ++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 526 insertions(+), 4 deletions(-) diff --git a/src/primitive_op.rs b/src/primitive_op.rs index eb5c317..b319307 100644 --- a/src/primitive_op.rs +++ b/src/primitive_op.rs @@ -40,17 +40,21 @@ use crate::ADKey; /// } /// /// impl PrimitiveOp for AddOp { +/// type ADContext = (); +/// /// fn add() -> Self { AddOp } /// fn linearize( /// &self, _b: &mut FragmentBuilder, /// _pi: &[GlobalValKey], _po: &[GlobalValKey], /// t: &[Option], +/// _ctx: &mut (), /// ) -> Vec> { /// vec![t[0].or(t[1])] /// } /// fn transpose_rule( /// &self, _b: &mut FragmentBuilder, /// ct: &[Option], _i: &[ValRef], _m: &OpMode, +/// _ctx: &mut (), /// ) -> Vec> { /// vec![ct[0], ct[0]] /// } @@ -60,6 +64,12 @@ pub trait PrimitiveOp: GraphOp where Self::InputKey: ADKey, { + /// Runtime AD context threaded through linearization and transpose. + /// + /// This can carry information such as concrete shapes or guard decisions + /// that influence how AD rules emit graph structure. + type ADContext: Default; + /// Returns the addition operation used for cotangent accumulation /// in `tidu::transpose`. When multiple cotangents flow to the same /// `GlobalValKey`, transpose emits `Op::add()` nodes to sum them. @@ -77,6 +87,7 @@ where primal_in: &[GlobalValKey], primal_out: &[GlobalValKey], tangent_in: &[Option], + ctx: &mut Self::ADContext, ) -> Vec> where Self: Sized; @@ -91,6 +102,7 @@ where cotangent_out: &[Option], inputs: &[ValRef], mode: &OpMode, + ctx: &mut Self::ADContext, ) -> Vec> where Self: Sized; diff --git a/tests/trait_tests.rs b/tests/trait_tests.rs index f420b7a..7aec62e 100644 --- a/tests/trait_tests.rs +++ b/tests/trait_tests.rs @@ -41,6 +41,8 @@ impl GraphOp for MockOp { } impl PrimitiveOp for MockOp { + type ADContext = (); + fn add() -> Self { MockOp::Add } @@ -51,6 +53,7 @@ impl PrimitiveOp for MockOp { primal_in: &[GlobalValKey], _primal_out: &[GlobalValKey], tangent_in: &[Option], + _ctx: &mut (), ) -> Vec> { match self { MockOp::Add => match (&tangent_in[0], &tangent_in[1]) { @@ -90,6 +93,7 @@ impl PrimitiveOp for MockOp { cotangent_out: &[Option], inputs: &[ValRef], _mode: &OpMode, + _ctx: &mut (), ) -> Vec> { match self { MockOp::Add => match &cotangent_out[0] { @@ -158,6 +162,7 @@ fn ad_key_higher_order_tangent() { #[test] fn primitive_op_linearize_add() { let mut builder = FragmentBuilder::::new(); + let mut ctx = (); let dx = builder.add_input(MockKey::User("dx".to_string())); let dy = builder.add_input(MockKey::User("dy".to_string())); @@ -168,7 +173,8 @@ fn primitive_op_linearize_add() { let primal_out = vec![GlobalValKey::Input(MockKey::User("sum".to_string()))]; let tangent_in = vec![Some(dx), Some(dy)]; - let result = MockOp::Add.linearize(&mut builder, &primal_in, &primal_out, &tangent_in); + let result = + MockOp::Add.linearize(&mut builder, &primal_in, &primal_out, &tangent_in, &mut ctx); assert_eq!(result.len(), 1); assert!(result[0].is_some()); @@ -186,6 +192,7 @@ fn primitive_op_linearize_add() { #[test] fn primitive_op_linearize_skip_inactive() { let mut builder = FragmentBuilder::::new(); + let mut ctx = (); let dx = builder.add_input(MockKey::User("dx".to_string())); let primal_in = vec![ @@ -195,7 +202,8 @@ fn primitive_op_linearize_skip_inactive() { let primal_out = vec![GlobalValKey::Input(MockKey::User("sum".to_string()))]; let tangent_in = vec![Some(dx), None]; - let result = MockOp::Add.linearize(&mut builder, &primal_in, &primal_out, &tangent_in); + let result = + MockOp::Add.linearize(&mut builder, &primal_in, &primal_out, &tangent_in, &mut ctx); assert_eq!(result.len(), 1); assert!(result[0].is_some()); @@ -206,6 +214,7 @@ fn primitive_op_linearize_skip_inactive() { #[test] fn primitive_op_transpose_add() { let mut builder = FragmentBuilder::::new(); + let mut ctx = (); let ct = builder.add_input(MockKey::User("ct".to_string())); let inputs = vec![ @@ -214,7 +223,13 @@ fn primitive_op_transpose_add() { ]; let cotangent_out = vec![Some(ct)]; - let result = MockOp::Add.transpose_rule(&mut builder, &cotangent_out, &inputs, &OpMode::Primal); + let result = MockOp::Add.transpose_rule( + &mut builder, + &cotangent_out, + &inputs, + &OpMode::Primal, + &mut ctx, + ); assert_eq!(result.len(), 2); assert_eq!(result[0], Some(ct)); @@ -226,6 +241,7 @@ fn primitive_op_transpose_add() { #[test] fn primitive_op_transpose_scale() { let mut builder = FragmentBuilder::::new(); + let mut ctx = (); let ct = builder.add_input(MockKey::User("ct".to_string())); let inputs = vec![ @@ -237,7 +253,8 @@ fn primitive_op_transpose_scale() { active_mask: vec![false, true], }; - let result = MockOp::Scale.transpose_rule(&mut builder, &cotangent_out, &inputs, &mode); + let result = + MockOp::Scale.transpose_rule(&mut builder, &cotangent_out, &inputs, &mode, &mut ctx); assert_eq!(result.len(), 2); assert!(result[0].is_none()); @@ -246,3 +263,496 @@ fn primitive_op_transpose_scale() { assert_eq!(frag.ops().len(), 1); assert_eq!(frag.ops()[0].op, MockOp::Scale); } + +#[derive(Default)] +struct RecordingContext { + linearize_calls: Vec, + transpose_calls: Vec, +} + +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +enum RecordingOp { + Add, + Foo, +} + +impl GraphOp for RecordingOp { + type Operand = f64; + type Context = (); + type InputKey = MockKey; + + fn n_inputs(&self) -> usize { + match self { + RecordingOp::Add => 2, + RecordingOp::Foo => 1, + } + } + + fn n_outputs(&self) -> usize { + 1 + } +} + +impl PrimitiveOp for RecordingOp { + type ADContext = RecordingContext; + + fn add() -> Self { + RecordingOp::Add + } + + fn linearize( + &self, + builder: &mut FragmentBuilder, + _primal_in: &[GlobalValKey], + _primal_out: &[GlobalValKey], + tangent_in: &[Option], + ctx: &mut RecordingContext, + ) -> Vec> { + ctx.linearize_calls.push( + match self { + RecordingOp::Add => "Add", + RecordingOp::Foo => "Foo", + } + .to_string(), + ); + + match self { + RecordingOp::Add => match (&tangent_in[0], &tangent_in[1]) { + (Some(dx), Some(dy)) => { + let out = builder.add_op( + RecordingOp::Add, + vec![ValRef::Local(*dx), ValRef::Local(*dy)], + OpMode::Linear { + active_mask: vec![true, true], + }, + ); + vec![Some(out[0])] + } + (Some(dx), None) => vec![Some(*dx)], + (None, Some(dy)) => vec![Some(*dy)], + (None, None) => vec![None], + }, + RecordingOp::Foo => vec![tangent_in.first().copied().flatten()], + } + } + + fn transpose_rule( + &self, + _builder: &mut FragmentBuilder, + cotangent_out: &[Option], + _inputs: &[ValRef], + _mode: &OpMode, + ctx: &mut RecordingContext, + ) -> Vec> { + ctx.transpose_calls.push( + match self { + RecordingOp::Add => "Add", + RecordingOp::Foo => "Foo", + } + .to_string(), + ); + + match self { + RecordingOp::Add => vec![cotangent_out[0], cotangent_out[0]], + RecordingOp::Foo => vec![cotangent_out[0]], + } + } +} + +#[derive(Default)] +struct BranchContext { + is_tall: bool, + guards: Vec, +} + +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +enum BranchOp { + Add, + Decompose, +} + +impl GraphOp for BranchOp { + type Operand = f64; + type Context = (); + type InputKey = MockKey; + + fn n_inputs(&self) -> usize { + match self { + BranchOp::Add => 2, + BranchOp::Decompose => 1, + } + } + + fn n_outputs(&self) -> usize { + 1 + } +} + +impl PrimitiveOp for BranchOp { + type ADContext = BranchContext; + + fn add() -> Self { + BranchOp::Add + } + + fn linearize( + &self, + builder: &mut FragmentBuilder, + _primal_in: &[GlobalValKey], + _primal_out: &[GlobalValKey], + tangent_in: &[Option], + ctx: &mut BranchContext, + ) -> Vec> { + match self { + BranchOp::Add => match (&tangent_in[0], &tangent_in[1]) { + (Some(dx), Some(dy)) => { + let out = builder.add_op( + BranchOp::Add, + vec![ValRef::Local(*dx), ValRef::Local(*dy)], + OpMode::Linear { + active_mask: vec![true, true], + }, + ); + vec![Some(out[0])] + } + (Some(dx), None) => vec![Some(*dx)], + (None, Some(dy)) => vec![Some(*dy)], + (None, None) => vec![None], + }, + BranchOp::Decompose => { + ctx.guards.push(ctx.is_tall); + match tangent_in[0] { + Some(dx) => { + let (op, inputs, active_mask) = if ctx.is_tall { + (BranchOp::Decompose, vec![ValRef::Local(dx)], vec![true]) + } else { + ( + BranchOp::Add, + vec![ValRef::Local(dx), ValRef::Local(dx)], + vec![true, true], + ) + }; + let out = builder.add_op(op, inputs, OpMode::Linear { active_mask }); + vec![Some(out[0])] + } + None => vec![None], + } + } + } + } + + fn transpose_rule( + &self, + builder: &mut FragmentBuilder, + cotangent_out: &[Option], + _inputs: &[ValRef], + _mode: &OpMode, + ctx: &mut BranchContext, + ) -> Vec> { + match self { + BranchOp::Add => match cotangent_out[0] { + Some(ct) => vec![Some(ct), Some(ct)], + None => vec![None, None], + }, + BranchOp::Decompose => { + ctx.guards.push(ctx.is_tall); + match cotangent_out[0] { + Some(ct) => { + let (op, inputs, active_mask) = if ctx.is_tall { + (BranchOp::Decompose, vec![ValRef::Local(ct)], vec![true]) + } else { + ( + BranchOp::Add, + vec![ValRef::Local(ct), ValRef::Local(ct)], + vec![true, true], + ) + }; + let out = builder.add_op(op, inputs, OpMode::Linear { active_mask }); + vec![Some(out[0])] + } + None => vec![None], + } + } + } + } +} + +#[test] +fn adcontext_linearize_records_calls() { + let mut builder = FragmentBuilder::::new(); + let mut ctx = RecordingContext::default(); + let dx = builder.add_input(MockKey::User("dx".to_string())); + + let primal_in = vec![GlobalValKey::Input(MockKey::User("x".to_string()))]; + let primal_out = vec![GlobalValKey::Input(MockKey::User("y".to_string()))]; + let tangent_in = vec![Some(dx)]; + + let result = + RecordingOp::Foo.linearize(&mut builder, &primal_in, &primal_out, &tangent_in, &mut ctx); + + assert_eq!(result, vec![Some(dx)]); + assert_eq!(ctx.linearize_calls, vec!["Foo".to_string()]); + assert!(ctx.transpose_calls.is_empty()); +} + +#[test] +fn adcontext_transpose_records_calls() { + let mut builder = FragmentBuilder::::new(); + let mut ctx = RecordingContext::default(); + let ct = builder.add_input(MockKey::User("ct".to_string())); + + let inputs = vec![ValRef::External(GlobalValKey::Input(MockKey::User( + "x".to_string(), + )))]; + let cotangent_out = vec![Some(ct)]; + + let result = RecordingOp::Foo.transpose_rule( + &mut builder, + &cotangent_out, + &inputs, + &OpMode::Linear { + active_mask: vec![true], + }, + &mut ctx, + ); + + assert_eq!(result, vec![Some(ct)]); + assert_eq!(ctx.transpose_calls, vec!["Foo".to_string()]); + assert!(ctx.linearize_calls.is_empty()); +} + +#[test] +fn adcontext_recording_context_accumulates_across_calls() { + let mut builder = FragmentBuilder::::new(); + let mut ctx = RecordingContext::default(); + let dx = builder.add_input(MockKey::User("dx".to_string())); + let ct = builder.add_input(MockKey::User("ct".to_string())); + + let primal_in = vec![GlobalValKey::Input(MockKey::User("x".to_string()))]; + let primal_out = vec![GlobalValKey::Input(MockKey::User("y".to_string()))]; + let tangent_in = vec![Some(dx)]; + let inputs = vec![ValRef::External(GlobalValKey::Input(MockKey::User( + "x".to_string(), + )))]; + let cotangent_out = vec![Some(ct)]; + + let linearized = + RecordingOp::Foo.linearize(&mut builder, &primal_in, &primal_out, &tangent_in, &mut ctx); + let transposed = RecordingOp::Foo.transpose_rule( + &mut builder, + &cotangent_out, + &inputs, + &OpMode::Linear { + active_mask: vec![true], + }, + &mut ctx, + ); + + assert_eq!(linearized, vec![Some(dx)]); + assert_eq!(transposed, vec![Some(ct)]); + assert_eq!(ctx.linearize_calls, vec!["Foo".to_string()]); + assert_eq!(ctx.transpose_calls, vec!["Foo".to_string()]); +} + +#[test] +fn adcontext_branch_tall_emits_decompose() { + let mut builder = FragmentBuilder::::new(); + let mut ctx = BranchContext { + is_tall: true, + ..Default::default() + }; + let dx = builder.add_input(MockKey::User("dx".to_string())); + + let primal_in = vec![GlobalValKey::Input(MockKey::User("x".to_string()))]; + let primal_out = vec![GlobalValKey::Input(MockKey::User("y".to_string()))]; + let tangent_in = vec![Some(dx)]; + + let result = + BranchOp::Decompose.linearize(&mut builder, &primal_in, &primal_out, &tangent_in, &mut ctx); + + assert_eq!(result.len(), 1); + assert!(result[0].is_some()); + assert_eq!(ctx.guards, vec![true]); + let frag = builder.build(); + assert_eq!(frag.ops().len(), 1); + assert_eq!(frag.ops()[0].op, BranchOp::Decompose); +} + +#[test] +fn adcontext_branch_wide_emits_add() { + let mut builder = FragmentBuilder::::new(); + let mut ctx = BranchContext { + is_tall: false, + ..Default::default() + }; + let dx = builder.add_input(MockKey::User("dx".to_string())); + + let primal_in = vec![GlobalValKey::Input(MockKey::User("x".to_string()))]; + let primal_out = vec![GlobalValKey::Input(MockKey::User("y".to_string()))]; + let tangent_in = vec![Some(dx)]; + + let result = + BranchOp::Decompose.linearize(&mut builder, &primal_in, &primal_out, &tangent_in, &mut ctx); + + assert_eq!(result.len(), 1); + assert!(result[0].is_some()); + assert_eq!(ctx.guards, vec![false]); + let frag = builder.build(); + assert_eq!(frag.ops().len(), 1); + assert_eq!(frag.ops()[0].op, BranchOp::Add); +} + +#[test] +fn adcontext_same_op_different_context_different_graph() { + let mut tall_builder = FragmentBuilder::::new(); + let mut wide_builder = FragmentBuilder::::new(); + let mut tall_ctx = BranchContext { + is_tall: true, + ..Default::default() + }; + let mut wide_ctx = BranchContext { + is_tall: false, + ..Default::default() + }; + let tall_dx = tall_builder.add_input(MockKey::User("dx_tall".to_string())); + let wide_dx = wide_builder.add_input(MockKey::User("dx_wide".to_string())); + + let primal_in = vec![GlobalValKey::Input(MockKey::User("x".to_string()))]; + let primal_out = vec![GlobalValKey::Input(MockKey::User("y".to_string()))]; + + BranchOp::Decompose.linearize( + &mut tall_builder, + &primal_in, + &primal_out, + &[Some(tall_dx)], + &mut tall_ctx, + ); + BranchOp::Decompose.linearize( + &mut wide_builder, + &primal_in, + &primal_out, + &[Some(wide_dx)], + &mut wide_ctx, + ); + + let tall_frag = tall_builder.build(); + let wide_frag = wide_builder.build(); + assert_eq!(tall_ctx.guards, vec![true]); + assert_eq!(wide_ctx.guards, vec![false]); + assert_eq!(tall_frag.ops().len(), 1); + assert_eq!(wide_frag.ops().len(), 1); + assert_eq!(tall_frag.ops()[0].op, BranchOp::Decompose); + assert_eq!(wide_frag.ops()[0].op, BranchOp::Add); + assert_ne!(tall_frag.ops()[0].op, wide_frag.ops()[0].op); +} + +#[test] +fn adcontext_transpose_branches_on_context() { + let mut tall_builder = FragmentBuilder::::new(); + let mut wide_builder = FragmentBuilder::::new(); + let mut tall_ctx = BranchContext { + is_tall: true, + ..Default::default() + }; + let mut wide_ctx = BranchContext { + is_tall: false, + ..Default::default() + }; + let tall_ct = tall_builder.add_input(MockKey::User("ct_tall".to_string())); + let wide_ct = wide_builder.add_input(MockKey::User("ct_wide".to_string())); + + let inputs = vec![ValRef::External(GlobalValKey::Input(MockKey::User( + "x".to_string(), + )))]; + let mode = OpMode::Linear { + active_mask: vec![true], + }; + + let tall_result = BranchOp::Decompose.transpose_rule( + &mut tall_builder, + &[Some(tall_ct)], + &inputs, + &mode, + &mut tall_ctx, + ); + let wide_result = BranchOp::Decompose.transpose_rule( + &mut wide_builder, + &[Some(wide_ct)], + &inputs, + &mode, + &mut wide_ctx, + ); + + assert_eq!(tall_result.len(), 1); + assert!(tall_result[0].is_some()); + assert_eq!(wide_result.len(), 1); + assert!(wide_result[0].is_some()); + assert_eq!(tall_ctx.guards, vec![true]); + assert_eq!(wide_ctx.guards, vec![false]); + let tall_frag = tall_builder.build(); + let wide_frag = wide_builder.build(); + assert_eq!(tall_frag.ops()[0].op, BranchOp::Decompose); + assert_eq!(wide_frag.ops()[0].op, BranchOp::Add); +} + +#[test] +fn adcontext_multiple_branch_points() { + let primal_in = vec![GlobalValKey::Input(MockKey::User("x".to_string()))]; + let primal_out = vec![GlobalValKey::Input(MockKey::User("y".to_string()))]; + + let mut tall_builder = FragmentBuilder::::new(); + let mut tall_ctx = BranchContext { + is_tall: true, + ..Default::default() + }; + let tall_dx = tall_builder.add_input(MockKey::User("dx_tall".to_string())); + let tall_first = BranchOp::Decompose.linearize( + &mut tall_builder, + &primal_in, + &primal_out, + &[Some(tall_dx)], + &mut tall_ctx, + ); + let tall_second = BranchOp::Decompose.linearize( + &mut tall_builder, + &primal_in, + &primal_out, + &[tall_first[0]], + &mut tall_ctx, + ); + + assert_eq!(tall_ctx.guards, vec![true, true]); + assert!(tall_second[0].is_some()); + let tall_frag = tall_builder.build(); + assert_eq!(tall_frag.ops().len(), 2); + assert_eq!(tall_frag.ops()[0].op, BranchOp::Decompose); + assert_eq!(tall_frag.ops()[1].op, BranchOp::Decompose); + + let mut wide_builder = FragmentBuilder::::new(); + let mut wide_ctx = BranchContext { + is_tall: false, + ..Default::default() + }; + let wide_dx = wide_builder.add_input(MockKey::User("dx_wide".to_string())); + let wide_first = BranchOp::Decompose.linearize( + &mut wide_builder, + &primal_in, + &primal_out, + &[Some(wide_dx)], + &mut wide_ctx, + ); + let wide_second = BranchOp::Decompose.linearize( + &mut wide_builder, + &primal_in, + &primal_out, + &[wide_first[0]], + &mut wide_ctx, + ); + + assert_eq!(wide_ctx.guards, vec![false, false]); + assert!(wide_second[0].is_some()); + let wide_frag = wide_builder.build(); + assert_eq!(wide_frag.ops().len(), 2); + assert_eq!(wide_frag.ops()[0].op, BranchOp::Add); + assert_eq!(wide_frag.ops()[1].op, BranchOp::Add); +} From 18a83fa6932ee05ff8ce392b9da1cffb8ef5e8d9 Mon Sep 17 00:00:00 2001 From: Hiroshi Shinaoka Date: Sat, 11 Apr 2026 21:11:26 +0900 Subject: [PATCH 2/2] ci: restore CI workflow deleted in v2 skeleton migration The workflow was accidentally removed in a2a35c7. Restores rustfmt, nextest (ubuntu/macos), and coverage jobs. Drops docs-site and check-coverage.py threshold check (scripts no longer exist). Co-Authored-By: Claude Opus 4.6 (1M context) --- .github/workflows/ci.yml | 61 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 61 insertions(+) create mode 100644 .github/workflows/ci.yml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..863bce9 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,61 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +permissions: + contents: read + +env: + CARGO_TERM_COLOR: always + +jobs: + fmt: + name: rustfmt + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: dtolnay/rust-toolchain@stable + with: + components: rustfmt + - uses: Swatinem/rust-cache@v2 + - run: cargo fmt --all --check + + nextest: + name: nextest (${{ matrix.os }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, macos-latest] + steps: + - uses: actions/checkout@v4 + - uses: dtolnay/rust-toolchain@stable + - uses: Swatinem/rust-cache@v2 + - uses: taiki-e/install-action@nextest + + - name: Run tests with nextest + run: cargo nextest run --workspace --release --no-fail-fast + + - name: Run doctests + run: cargo test --doc --workspace --release + + coverage: + name: coverage + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: dtolnay/rust-toolchain@stable + with: + components: llvm-tools-preview + - uses: Swatinem/rust-cache@v2 + - uses: taiki-e/install-action@nextest + - uses: taiki-e/install-action@cargo-llvm-cov + - name: Generate coverage report + run: cargo llvm-cov nextest --workspace --release