Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion lib/ModelingToolkitTearing/Project.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
name = "ModelingToolkitTearing"
uuid = "6bb917b9-1269-42b9-9f7c-b0dca72083ab"
version = "1.20.6"
version = "1.20.7"
authors = ["Aayush Sabharwal <aayush.sabharwal@juliahub.com>"]

[deps]
Expand Down
11 changes: 11 additions & 0 deletions lib/ModelingToolkitTearing/src/tearingstate.jl
Original file line number Diff line number Diff line change
Expand Up @@ -1072,6 +1072,17 @@ function shift_discrete_system(ts::TearingState)
# necessary for hybrid system handling.
(; fullvars, sys) = ts
fullvars_set = Set{SymbolicT}(fullvars)
# `fullvars` only contains scalarized array elements, so an array variable used
# without scalarization (as in `f(x(k - 1))` for an array `x`) is absent from
# `fullvars_set` and would be left unshifted while its elements are shifted, which
# puts the two forms one tick apart. Add the arrays the scalarized elements belong
# to, so that both forms are shifted alike; the `isoperator` check below discards
# the entries that are a `Sample`, `Hold` or `Pre` of an array.
for v in fullvars
arr, isidx = MTKBase.split_indexed_var(v)
isidx || continue
push!(fullvars_set, strip_shifts(arr))
end
discvars = OrderedSet{SymbolicT}()
eqs = equations(sys)
for eq in eqs
Expand Down
12 changes: 12 additions & 0 deletions lib/ModelingToolkitTearing/src/utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,18 @@ function descend_lower_shift_varname(var, iv)
end
end

"""
$TYPEDSIGNATURES

Remove all `Shift`s applied to `var` and return the variable they are applied to.
"""
function strip_shifts(var::SymbolicT)
@match var begin
BSImpl.Term(; f, args) && if f isa Shift end => strip_shifts(args[1])
_ => var
end
end

function is_time_dependent_parameter(p::SymbolicT, allps::Set{SymbolicT}, iv::SymbolicT)
return p in allps && @match p begin
BSImpl.Term(; f, args) => begin
Expand Down
59 changes: 59 additions & 0 deletions lib/ModelingToolkitTearing/test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -735,3 +735,62 @@ end
@test prob[y] == 3.0
@test prob[z] == 3.0
end

# A function of a whole array which cannot be scalarized, so that the unscalarized form of
# an array variable survives into the compiled system.
combine_array(v) = sum(v)
@register_symbolic combine_array(v::AbstractVector)

@testset "Shifts applied to whole array variables" begin
# https://github.com/SciML/ModelingToolkit.jl/issues/5169
k = ShiftIndex(t)
@variables y(t) ud(t) (xd(t))[1:2]
# The elements of `xd` are aliased away, so `fullvars` only holds the scalarized
# elements while the equation for `ud` uses the whole array.
common_eqs = [y ~ y(k - 1) + 1, xd[1] ~ y, xd[2] ~ 2y]

@named sys = System([common_eqs; ud ~ combine_array(xd(k - 1))], t)
@named ref = System([common_eqs; ud ~ xd(k - 1)[1] + xd(k - 1)[2]], t)
ss = mtkcompile(sys)
ssref = mtkcompile(ref)

@test issetequal(unknowns(ss), unknowns(ssref))
udeq = only(filter(eq -> isequal(eq.lhs, unwrap(ud)), observed(ss)))
# `xd` is used one tick back, so the observed equation for `ud` must refer to the
# same shifted variable the unknowns are lowered to.
prev = only(arguments(udeq.rhs))
@test issubset(Set(collect(prev)), Set(unknowns(ss)))

u0 = [y(k - 1) => 0.0, xd(k - 1) => [3.0, 4.0]]
prob = DiscreteProblem(ss, u0, (0, 5))
probref = DiscreteProblem(ssref, u0, (0, 5))
@test prob[ud] == probref[ud]
end

@testset "Shifts applied to whole array variables in a clock partition" begin
# https://github.com/SciML/ModelingToolkit.jl/issues/5169
dt = 0.1
k = ShiftIndex(Clock(dt))
@variables x(t) y(t) u(t) ud(t) (xd(t))[1:2]
eqs = [
xd[1] ~ Sample(dt)(y)
xd[2] ~ 2 * Sample(dt)(y)
ud ~ combine_array(xd(k - 1))
u ~ Hold(ud)
D(x) ~ -x + u
y ~ x
]
@named sys = System(eqs, t)
ci = MTKTearing.infer_clocks!(MTKTearing.ClockInference(TearingState(sys)))
tss, _, continuous_id, _ = MTKTearing.split_system(ci)
disc = tss[findfirst(!=(continuous_id), eachindex(tss))]

# `split_system` shifts the discrete partition forward by one tick, and the
# unscalarized `xd` must be shifted along with its elements.
vars = Set{Symbolics.SymbolicT}()
for eq in equations(disc)
SU.search_variables!(vars, eq; is_atomic = MTKBase.OperatorIsAtomic{SU.Operator}())
end
@test unwrap(xd) in vars
@test !(unwrap(xd(k - 1)) in vars)
end
Loading