From 9a67fa6386ece70dd4ff498f82a81c6804af9a2b Mon Sep 17 00:00:00 2001 From: Jukka Aho Date: Sat, 22 Nov 2025 12:07:11 +0200 Subject: [PATCH] feat(basis): Generate nbasis functions with canonical type names MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Extend basis_generator.jl to automatically emit nbasis() functions alongside get_basis_functions() and get_basis_derivatives(): 1. Add canonical_type_name() helper to convert type aliases to full forms (Quad9 → Quadrilateral{9}, Tri3 → Triangle{3}) 2. Generate @inline nbasis(::Topology, ::Basis) = N for each basis family This provides zero-cost compile-time basis function count 3. Use instance-based dispatch (::Triangle{3}, not ::Type{Triangle{3}}) for consistency with get_basis_* functions 4. Enhanced pretty-printer to handle both function definitions and simple assignments, properly unwrapping begin...end blocks Generated code now uses canonical topology names throughout, avoiding historical type aliases for better clarity and API consistency. --- src/basis/basis_generator.jl | 69 +++++++++++++++++++++++++++++++++--- 1 file changed, 64 insertions(+), 5 deletions(-) diff --git a/src/basis/basis_generator.jl b/src/basis/basis_generator.jl index edae1af..794827a 100644 --- a/src/basis/basis_generator.jl +++ b/src/basis/basis_generator.jl @@ -441,11 +441,30 @@ end # Code generation # ------------------------------------------------------------------------------ +""" + canonical_type_name(T::Type) -> String + +Convert a type to its canonical string representation, avoiding type aliases. +For example: Quad9 → "Quadrilateral{9}", Tri3 → "Triangle{3}" +""" +function canonical_type_name(T::Type) + # Get the type name and parameters + name = nameof(T.name.wrapper) + params = T.parameters + + if isempty(params) + return string(name) + else + param_strs = [isa(p, Type) ? canonical_type_name(p) : string(p) for p in params] + return string(name, "{", join(param_strs, ","), "}") + end +end + function create_basis_code(topology_type_expr, basis_type_expr::Type, description, X::Vector{<:Vecish{D}}, ansatz::Expr) where D V = vandermonde_matrix(ansatz, X) basis = calculate_interpolation_polynomials(ansatz, V) dbasis = calculate_interpolation_polynomial_derivatives(basis, D) - N = length(X) + N = length(X) # ← This is the number of basis functions! # Apply numerical filtering to remove noise from floating point arithmetic basis = [filter_expr(b) for b in basis] @@ -478,15 +497,26 @@ function create_basis_code(topology_type_expr, basis_type_expr::Type, descriptio func_body = Base.remove_linenums!(func_body) dfunc_body = Base.remove_linenums!(dfunc_body) + # Get canonical type names (avoiding aliases like Quad9 → Quadrilateral{9}) + topology_canon = canonical_type_name(topology_type_expr) + basis_canon = canonical_type_name(basis_type_expr) + + # Build type expressions with canonical names + topology_expr = Meta.parse(topology_canon) + basis_expr = Meta.parse(basis_canon) + # Splat the body contents directly (avoid nested begin/end) return quote - @inline function get_basis_functions(::$topology_type_expr, ::$basis_type_expr, xi::Vec{$D,T}) where T + @inline function get_basis_functions(::$topology_expr, ::$basis_expr, xi::Vec{$D,T}) where T $(func_body.args...) end - @inline function get_basis_derivatives(::$topology_type_expr, ::$basis_type_expr, xi::Vec{$D,T}) where T + @inline function get_basis_derivatives(::$topology_expr, ::$basis_expr, xi::Vec{$D,T}) where T $(dfunc_body.args...) end + + # Number of basis functions (compile-time constant for validation) + @inline nbasis(::$topology_expr, ::$basis_expr) = $N end end @@ -522,6 +552,9 @@ function write_generated_file(elements) println(output, "# Why SVector? Enables natural vector operations:") println(output, "# u_interp = dot(node_values, N)") println(output, "# grad_u = sum(node_values[i] * dN[i] for i in 1:N)") + println(output, "#") + println(output, "# Also generates: nbasis(basis_type, topology_type) -> Int") + println(output, "# Zero-cost function returning number of basis functions (compile-time constant)") for (i, elem) in enumerate(elements) println("[$i/$(length(elements))] Generating $(elem.name)...") @@ -540,9 +573,35 @@ function write_generated_file(elements) continue end - # Extract function components + # Handle @inline macro expressions if expr isa Expr && expr.head == :macrocall && expr.args[1] == Symbol("@inline") - func = expr.args[3] + inner_expr = expr.args[3] + + # Case 1: Simple assignment @inline nbasis(...) = N or @inline nbasis(...) = begin N end + if inner_expr isa Expr && inner_expr.head == :(=) + lhs = inner_expr.args[1] + rhs = inner_expr.args[2] + + # Unwrap begin...end block if present + if rhs isa Expr && rhs.head == :block + # Extract the actual value (skip line numbers) + actual_value = nothing + for arg in rhs.args + if !(arg isa LineNumberNode) + actual_value = arg + break + end + end + rhs = actual_value + end + + println(output, "@inline ", lhs, " = ", rhs) + println(output) + continue + end + + # Case 2: Function definition @inline function f(...) ... end + func = inner_expr func_sig = func.args[1] func_body = func.args[2]