From d81dcd54ee2a8b3689e8e47907333f0293d819c0 Mon Sep 17 00:00:00 2001 From: lkdvos Date: Wed, 12 Aug 2026 16:07:01 +0200 Subject: [PATCH] fix `@planar` backend and allocator insertion `_insert_planar_operations` rewrites the tensor calls to `GlobalRef(TensorKit, :planar*!)`, which `TO.insertargument` does not match, so `@planar backend=... ` and `@planar allocator=...` silently dropped both arguments from the planar calls. Co-Authored-By: Claude Opus 5 (1M context) --- src/planar/macros.jl | 6 ++++-- src/planar/postprocessors.jl | 16 ++++++++++++++-- test/tensors/planar.jl | 27 +++++++++++++++++++++++++++ 3 files changed, 45 insertions(+), 4 deletions(-) diff --git a/src/planar/macros.jl b/src/planar/macros.jl index c2afcb596..16f7f1c0a 100644 --- a/src/planar/macros.jl +++ b/src/planar/macros.jl @@ -30,7 +30,7 @@ function planarparser(planarexpr, kwargs...) if name == :backend hasbackend = true backend = val - push!(parser.postprocessors, ex -> TO.insertbackend(ex, backend)) + push!(parser.postprocessors, ex -> insertplanarbackend(ex, backend)) break end end @@ -39,8 +39,10 @@ function planarparser(planarexpr, kwargs...) allocator = val if !hasbackend backend = Expr(:call, GlobalRef(TensorOperations, :DefaultBackend)) - push!(parser.postprocessors, ex -> TO.insertbackend(ex, backend)) + push!(parser.postprocessors, ex -> insertplanarbackend(ex, backend)) end + push!(parser.postprocessors, ex -> insertplanarallocator(ex, allocator)) + # the alloc/free calls are still `GlobalRef(TensorOperations, ...)` push!(parser.postprocessors, ex -> TO.insertallocator(ex, allocator)) break end diff --git a/src/planar/postprocessors.jl b/src/planar/postprocessors.jl index 7ae527383..8cf55ea2f 100644 --- a/src/planar/postprocessors.jl +++ b/src/planar/postprocessors.jl @@ -90,6 +90,18 @@ function _insert_planar_operations(ex) return ex end +# like `TO.insertargument`, but matching `GlobalRef`s into `TensorKit` +function _insertargument(ex, arg, methods) + if isexpr(ex, :call) && ex.args[1] isa GlobalRef && + ex.args[1].mod === TensorKit && ex.args[1].name ∈ methods + return Expr(:call, ex.args..., arg) + elseif isa(ex, Expr) + return Expr(ex.head, (_insertargument(e, arg, methods) for e in ex.args)...) + else + return ex + end +end + """ insertplanarbackend(ex, backend) @@ -98,7 +110,7 @@ Insert the backend argument into the tensor operation methods `planaradd!`, `pla See also: [`TensorOperations.insertbackend`](@ref). """ function insertplanarbackend(ex, backend) - return TO.insertargument(ex, backend, (:planaradd!, :planartrace!, :planarcontract!)) + return _insertargument(ex, backend, _PLANAR_OPERATIONS) end """ @@ -109,5 +121,5 @@ Insert the allocator argument into the tensor operation methods `planaradd!`, `p See also: [`TensorOperations.insertallocator`](@ref). """ function insertplanarallocator(ex, allocator) - return TO.insertargument(ex, allocator, (:planaradd!, :planartrace!, :planarcontract!)) + return _insertargument(ex, allocator, _PLANAR_OPERATIONS) end diff --git a/test/tensors/planar.jl b/test/tensors/planar.jl index 8e8a3e778..99ae49ebd 100644 --- a/test/tensors/planar.jl +++ b/test/tensors/planar.jl @@ -101,6 +101,33 @@ end @testset "@planar" verbose = true begin T = ComplexF64 + @testset "backend and allocator insertion" begin + # trailing arguments of every `planar*!` call in `ex` + function planartrailing(ex, out = Any[]) + ex isa Expr || return out + if Meta.isexpr(ex, :call) && ex.args[1] isa GlobalRef && + ex.args[1].name in (:planaradd!, :planartrace!, :planarcontract!) + push!(out, ex.args[end]) + end + foreach(a -> planartrailing(a, out), ex.args) + return out + end + + ex = @macroexpand @planar backend = MarkerBackend() C[i; j] := A[i; k l] * + τ[k l; m n] * B[m n; j] + trailing = planartrailing(ex) + @test !isempty(trailing) + @test all(==(:(MarkerBackend())), trailing) + + # an allocator implies a default backend, and both land on the planar calls + ex = @macroexpand @planar allocator = MarkerAllocator() C[i; j] := A[i; k l] * + τ[k l; m n] * B[m n; j] + trailing = planartrailing(ex) + @test !isempty(trailing) + @test all(==(:(MarkerAllocator())), trailing) + @test occursin("DefaultBackend", string(ex)) + end + @testset "contractcheck" begin V = ℂ^2 A = rand(T, V ⊗ V ← V)