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
61 changes: 19 additions & 42 deletions src/Enzyme.jl
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@ function assemble_scalar_enzyme_safe!(
U_old = p.field_old
_update_for_assembly!(p, dof, Uu)
conns = fspace.elem_conns
# foreach_block(fspace, p) do physics, props, ref_fe, b
for (b, (
block_physics, ref_fe
)) in enumerate(zip(
Expand All @@ -20,64 +19,49 @@ function assemble_scalar_enzyme_safe!(
_assemble_scalar_block_enzyme_safe!(
KA.CPU(),
block_view(storage, b),
conns.data, conns.offsets[b],
conns,
func,
b,
block_physics, ref_fe,
X, t, Δt,
U, U_old,
block_view(p.state_old, b), block_view(p.state_new, b),
p.properties,
p.state_old, p.state_new, p.properties,
return_type
)
end
end

function _assemble_scalar_block_enzyme_safe!(
::KA.CPU,
# field::AbstractField,
field,
conns::Conn, coffset::Int,
conns_all,
func::Function,
b::Int,
physics::AbstractPhysics, ref_fe::ReferenceFE,
X::AbstractField, t::T, dt::T,
U::Solution, U_old::Solution,
state_old::S, state_new::S, props::AbstractArray,
state_old::StateVariableField, state_new::StateVariableField, props::PropertyField,
return_type::R
) where {
T <: Number,
Conn <: AbstractArray,
Solution <: AbstractField,
S, #<: L2QuadratureField
R <: AssembledReturnType
}

for e in axes(state_old, 3)
conns = conns_all.data
coffset = conns_all.offsets[b]
for e in 1:conns_all.nelems[b]
conn = connectivity(ref_fe, conns, e, coffset)
x_el, u_el, u_el_old = element_level_fields(ref_fe, conn, X, U, U_old)
props_el = properties(props, e, b)
# val_el = _element_scratch(return_type, ref_fe, U)

for q in 1:num_cell_quadrature_points(ref_fe)
interps = _cell_interpolants(ref_fe, q)
state_old_q = _quadrature_level_state(state_old, q, e)
state_new_q = _quadrature_level_state(state_new, q, e)
state_old_q = state_variables(state_old, q, e, b)
state_new_q = state_variables(state_new, q, e, b)
val_q = func(physics, interps, x_el, t, dt, u_el, u_el_old, state_old_q, state_new_q, props_el)
# val_el = _accumulate_q_value(return_type, field, val_q, val_el, q, e)
field[1, q, e] = val_q
end
# _assemble_element!(field, val_el, conn, e)

# writing inline to avoid atomic call
# n_dofs = size(field, 1)
# for d in axes(field, 1)
# for n in axes(conn, 1)
# global_id = n_dofs * (conn[n] - 1) + d
# local_id = n_dofs * (n - 1) + d
# field.data[global_id] += val_el[local_id]
# end
# end
end
return nothing
end
Expand Down Expand Up @@ -118,18 +102,15 @@ function assemble_vector_enzyme_safe!(
values(p.physics), values(fspace.ref_fes)
))
_assemble_vector_block_enzyme_safe!(
# KA.get_backend(storage),
KA.CPU(),
storage,
conns.data, conns.offsets[b],
conns,
func,
b,
block_physics, ref_fe,
X, t, Δt,
U, U_old,
block_view(p.state_old, b), block_view(p.state_new, b),
p.properties
# return_type
p.state_old, p.state_new, p.properties
)
end

Expand All @@ -146,25 +127,21 @@ TODO add state variables and physics properties
"""
function _assemble_vector_block_enzyme_safe!(
::KA.CPU,
# field::AbstractField,
field,
conns::Conn, coffset::Int,
conns_all,
func::Function,
b::Int,
physics::AbstractPhysics, ref_fe::ReferenceFE,
X::AbstractField, t::T, dt::T,
U::Solution, U_old::Solution,
state_old::S, state_new::S, props::AbstractArray,
# return_type::R
state_old::StateVariableField, state_new::StateVariableField, props::PropertyField,
) where {
T <: Number,
Conn <: AbstractArray,
Solution <: AbstractField,
S, #<: L2QuadratureField
# R <: AssembledReturnType
Solution <: AbstractField
}

for e in axes(state_old, 3)
conns = conns_all.data
coffset = conns_all.offsets[b]
for e in 1:conns_all.nelems[b]
conn = connectivity(ref_fe, conns, e, coffset)
x_el, u_el, u_el_old = element_level_fields(ref_fe, conn, X, U, U_old)

Expand All @@ -173,8 +150,8 @@ function _assemble_vector_block_enzyme_safe!(

for q in 1:num_cell_quadrature_points(ref_fe)
interps = _cell_interpolants(ref_fe, q)
state_old_q = _quadrature_level_state(state_old, q, e)
state_new_q = _quadrature_level_state(state_new, q, e)
state_old_q = state_variables(state_old, q, e, b)
state_new_q = state_variables(state_new, q, e, b)
# val_q = func(physics, interps, x_el, t, dt, u_el, u_el_old, state_old_q, state_new_q, props_el)
# val_el = _accumulate_q_value(return_type, field, val_q, val_el, q, e)

Expand Down
73 changes: 55 additions & 18 deletions src/Fields.jl
Original file line number Diff line number Diff line change
Expand Up @@ -548,11 +548,17 @@ function num_fields(field::PropertyField, b::Int)
end

function properties(field::PropertyField, e::Int, b::Int)
@assert 1 <= b && b <= field.nblocks
offset = field.offsets[b]
nfields = num_fields(field, b)
# `if/elseif` with no `else` leaves `start` undefined on any third value,
# which surfaces as an `UndefVarError` rather than saying what went wrong.
# Only two layouts exist, so make the second branch total and assert it.
if field.isblockconstant[b] == PROPS_CONST
start = offset
elseif field.isblockconstant[b] == PROPS_ELEMS
else
@assert field.isblockconstant[b] == PROPS_ELEMS
@assert 1 <= e && e <= field.nelems[b]
start = offset + nfields * (e - 1)
end
return PropertyFieldView(field.data, start, nfields)
Expand All @@ -564,30 +570,31 @@ struct PropertyFieldView{T, D <: AbstractVector{T}} <: AbstractVector{T}
len::Int
end

Base.size(v::PropertyFieldView) = (v.len,)
Base.length(v::PropertyFieldView) = v.len
Base.IndexStyle(::Type{<:PropertyFieldView}) = IndexLinear()
Base.@propagate_inbounds function Base.getindex(v::PropertyFieldView, i::Int)
@boundscheck checkbounds(v, i)
return @inbounds v.data[v.start + i - 1]
end
Base.IndexStyle(::Type{<:PropertyFieldView}) = IndexLinear()
Base.length(v::PropertyFieldView) = v.len
Base.size(v::PropertyFieldView) = (v.len,)

######################################################################################################
# StateVariableField
######################################################################################################
struct StateVariableField{
T, # Let it be anything to allow for structs
D <: AbstractVector{T}
D <: AbstractVector{T},
I <: AbstractVector{Int}
} <: AbstractDiscontinuousField{T, D}
data::D # flat storage (CPU or GPU)
nblocks::Int
nfields::Vector{Int}
nepes::Vector{Int} # num nodes, q points, etc.
nelems::Vector{Int}
offsets::Vector{Int}
nfields::I
nepes::I # num nodes, q points, etc.
nelems::I
offsets::I

function StateVariableField{T, D}(data, nblocks, nfields, nepes, nelems, offsets) where {T, D}
new{T, D}(data, nblocks, nfields, nepes, nelems, offsets)
function StateVariableField{T, D, I}(data, nblocks, nfields, nepes, nelems, offsets) where {T, D, I}
new{T, D, I}(data, nblocks, nfields, nepes, nelems, offsets)
end

function StateVariableField(arrs::Vector{<:AbstractArray{T, 3}}) where T
Expand All @@ -601,7 +608,7 @@ struct StateVariableField{
offset += nfields[b] * nepes[b] * nelems[b]
end
data = mapreduce(vec, vcat, arrs)
return StateVariableField{T, typeof(data)}(data, length(nepes), nfields, nepes, nelems, offsets)
return StateVariableField{T, typeof(data), typeof(nepes)}(data, length(nepes), nfields, nepes, nelems, offsets)
end

function StateVariableField(::UndefInitializer, ::Type{T}, nfields::Int, qsizes::Vector{Tuple{Int, Int}}) where T
Expand All @@ -621,15 +628,16 @@ struct StateVariableField{
end
end

function Adapt.adapt_structure(to, field::StateVariableField{T, D}) where {T, D}
function Adapt.adapt_structure(to, field::StateVariableField{T, D, I}) where {T, D, I}
data = adapt(to, field.data)
return StateVariableField{T, typeof(data)}(
nfields = adapt(to, field.nfields)
return StateVariableField{T, typeof(data), typeof(nfields)}(
data,
field.nblocks,
field.nfields,
field.nepes,
field.nelems,
field.offsets
nfields,
adapt(to, field.nepes),
adapt(to, field.nelems),
adapt(to, field.offsets)
)
end

Expand All @@ -648,3 +656,32 @@ end
function num_fields(field::StateVariableField, b::Int)
return field.nfields[b]
end

function state_variables(field::StateVariableField, q::Int, e::Int, b::Int)
@assert 1 <= q <= field.nepes[b]
@assert 1 <= e <= field.nelems[b]
offset = field.offsets[b]
nfields = field.nfields[b]
nqs = field.nepes[b]
start = offset + nfields * (q - 1) + nfields * nqs * (e - 1)
return StateVariableFieldView(field.data, start, nfields)
end

struct StateVariableFieldView{T, D <: AbstractVector{T}} <: AbstractVector{T}
data::D
start::Int
len::Int
end

Base.@propagate_inbounds function Base.getindex(v::StateVariableFieldView, i::Int)
@boundscheck checkbounds(v, i)
return @inbounds v.data[v.start + i - 1]
end
Base.IndexStyle(::Type{<:StateVariableFieldView}) = IndexLinear()
Base.length(v::StateVariableFieldView) = v.len
Base.@propagate_inbounds function Base.setindex!(v::StateVariableFieldView{T, D}, val::T, i::Int) where {T, D}
@boundscheck checkbounds(v, i)
@inbounds v.data[v.start + i - 1] = val
return nothing
end
Base.size(v::StateVariableFieldView) = (v.len,)
26 changes: 26 additions & 0 deletions src/Utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -267,6 +267,32 @@ end
quote $(exprs...) end
end

"""
takes in connectivity and element id
"""
function foreach_element(
f, conns, block_id,
backend = KA.get_backend(conns.data);
max_tasks = Threads.nthreads(),
min_elems = 1,
prefer_threads::Bool = true,
# GPU settings
block_size = 256
)
nelems = conns.nelems[block_id]
if AK.use_gpu_algorithm(backend, prefer_threads)
AK._forindices_gpu(f, 1:nelems, backend; block_size)
elseif max_tasks == 1
_forindices_serial(f, 1:nelems)
else
_forindices_threads(
f, 1:nelems,
max_tasks = max_tasks,
min_elems = min_elems
)
end
end

#########################################
# hooks for extensions
#########################################
Expand Down
Loading
Loading