Skip to content

Commit 7b49315

Browse files
Avoid eval in PDEBPINN differential masking (#1087)
Co-authored-by: ChrisRackauckas-Claude <accounts@chrisrackauckas.com>
1 parent 091a3e9 commit 7b49315

1 file changed

Lines changed: 18 additions & 19 deletions

File tree

test/PDEBPINN/bpinn_pde__bpinn_pde_inv_iii_improved_parametric_kuromo_sivashinsky_equation_solve.jl

Lines changed: 18 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -4,28 +4,25 @@ using Test
44
@testset "BPINN PDE Inv III: Improved Parametric Kuromo-Sivashinsky Equation solve" begin
55
using MCMCChains, Lux, ModelingToolkit, Distributions, OrdinaryDiffEq,
66
AdvancedHMC, LogDensityProblems, Statistics, Random, Functors, NeuralPDE, MonteCarloMeasurements,
7-
ComponentArrays
7+
ComponentArrays, SymbolicUtils
88
import DomainSets: Interval, infimum, supremum
99

1010
Random.seed!(100)
1111

12-
function recur_expression(exp, Dict_differentials)
13-
for in_exp in exp.args
14-
if !(in_exp isa Expr)
15-
# skip +,== symbols, characters etc
16-
continue
17-
18-
elseif in_exp.args[1] isa ModelingToolkit.Differential
19-
# first symbol of differential term
20-
# Dict_differentials for masking differential terms
21-
# and resubstituting differentials in equations after putting in interpolations
22-
# temp = in_exp.args[end]
23-
Dict_differentials[eval(in_exp)] = Symbolics.variable("diff_$(length(Dict_differentials) + 1)")
24-
return
25-
else
26-
recur_expression(in_exp, Dict_differentials)
27-
end
12+
function recur_expression(term, Dict_differentials)
13+
term = SymbolicUtils.unwrap(term)
14+
SymbolicUtils.iscall(term) || return nothing
15+
16+
op = SymbolicUtils.operation(term)
17+
if op isa ModelingToolkit.Differential
18+
Dict_differentials[term] = Symbolics.variable("diff_$(length(Dict_differentials) + 1)")
19+
return nothing
2820
end
21+
22+
for arg in SymbolicUtils.arguments(term)
23+
recur_expression(arg, Dict_differentials)
24+
end
25+
return nothing
2926
end
3027

3128
@parameters α
@@ -122,8 +119,10 @@ using Test
122119
# neccesarry for loss function construction (involves Operator masking)
123120
eqs = pde_system.eqs
124121
Dict_differentials = Dict()
125-
exps = toexpr.(eqs)
126-
nullobj = [recur_expression(exp, Dict_differentials) for exp in exps]
122+
for eq in eqs
123+
recur_expression(eq.lhs, Dict_differentials)
124+
recur_expression(eq.rhs, Dict_differentials)
125+
end
127126

128127
# Dict_differentials is now ;
129128
# Dict{Any, Any} with 5 entries:

0 commit comments

Comments
 (0)