Skip to content
1 change: 1 addition & 0 deletions LICENSE
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
MIT License

Copyright (c) 2023 Julian Trommer
Copyright (c) 2026 Josef Kircher

Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
Expand Down
32 changes: 19 additions & 13 deletions Project.toml
Original file line number Diff line number Diff line change
@@ -1,45 +1,51 @@
name = "GraphNetCore"
uuid = "7809f980-de1b-4f9a-8451-85f041491431"
authors = ["JT <julian.trommer@uni-a.de>"]
version = "0.3.1"
authors = ["Julian Trommer <julian.trommer@uni-a.de>"]
version = "0.4.0"

[deps]
Adapt = "79e6a3ab-5dfb-504d-930d-738a2a938a0e"
CUDA = "052768ef-5323-5732-b1bb-66c8b64840ba"
ComponentArrays = "b0b7db55-cfe3-40fc-9ded-d10e2dbeff66"
DataFrames = "a93c6f00-e57d-5684-b7b6-d8193f3e46c0"
ForwardDiff = "f6369f11-7733-5829-9624-2563aa707210"
Enzyme = "7da242da-08ed-463a-9acd-ee780be4f1d9"
JLD2 = "033835bb-8acc-5ee8-8aae-3f567f8a3819"
KernelAbstractions = "63c18a36-062a-441e-b654-da1e3ab1ce7c"
Lux = "b2108857-7c20-44ae-9111-449ecde12c47"
LuxCUDA = "d0bbae9a-e099-4d5b-a835-1c6931763bda"
NNlib = "872c559c-99b0-510c-b3b7-b6c96a88d5cd"
Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c"
Reactant = "3c362404-f566-11ee-1572-e11a4b42c853"
Setfield = "efcf1570-3423-57d1-acb7-fd33fddbac46"
Statistics = "10745b16-79ce-11e8-11f9-7d13ad32a3b2"
Tullio = "bc48ee85-29a4-5162-ae0b-a64e1601d4bc"
Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f"
cuDNN = "02a925ec-e4fe-4b08-9a7e-0d78e3d38ccd"

[compat]
Adapt = "4.2"
Aqua = "0.8"
CUDA = "5"
CUDA = "5.9.6"
ComponentArrays = "0.15"
DataFrames = "1.6"
ForwardDiff = "0.10"
JLD2 = "0.4"
Enzyme = "0.13.73"
JLD2 = "0.6"
KernelAbstractions = "0.9"
Lux = "0.5"
LuxCUDA = "0.3"
Lux = "1.13"
NNlib = "0.9"
Random = "1"
Reactant = "0.2.169"
Setfield = "1.1.2"
Statistics = "1"
Tullio = "0.3.7"
Zygote = "0.6"
cuDNN = "1.3"
Test = "1"
julia = "1.10"
Tullio = "0.3.7"
Zygote = "0.6, 0.7"
cuDNN = "1.4.5"
julia = "1.11"

[extras]
Aqua = "4c88cf16-eb10-579e-8560-4a9242c79595"
CUDA_Runtime_jll = "76a88914-d11a-5bdc-97e0-2f5a05c973a2"
Lux = "b2108857-7c20-44ae-9111-449ecde12c47"
Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40"

[targets]
Expand Down
7 changes: 4 additions & 3 deletions src/GraphNetCore.jl
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,9 @@

module GraphNetCore

using CUDA
using Lux, LuxCUDA
using Lux
using CUDA, cuDNN
using Reactant
using Tullio
using Random

Expand All @@ -23,7 +24,7 @@ export NormaliserOffline, NormaliserOfflineMinMax, NormaliserOfflineMeanStd,
NormaliserOnline

# graph_network.jl
export build_model, step!, save!, load
export build_model, step!, set_training!, save!, load
# normaliser.jl
export inverse_data
# utils.jl
Expand Down
34 changes: 10 additions & 24 deletions src/feature_graph.jl
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
# Licensed under the MIT license. See LICENSE file in the project root for details.
#

using NNlib

"""
FeatureGraph(nf, ef, senders, receivers)

Expand All @@ -14,33 +16,16 @@ Data structure that is used as an input for the [`GraphNetwork`](@ref).
- `senders`: List of nodes in the mesh where graph edges start.
- `receivers`: List of nodes in the mesh where graph edges end.
"""
mutable struct FeatureGraph{F <: AbstractArray, T <: AbstractArray}
nf::F
ef::F
mutable struct FeatureGraph{T <: AbstractArray}
nf::Any
ef::Any
senders::T
receivers::T
end

"""
update_features!(g; nf, ef)

Updates the node and edge features of the given [`FeatureGraph`](@ref).

## Arguments
- `g`: [`FeatureGraph`](@ref) that should be updated.

## Keyword Arguments
- `nf`: Updated node features.
- `ef`: Updated edge features.

## Returns
- Updated graph as a [`FeatureGraph`](@ref) struct.
"""
function update_features!(g::FeatureGraph; nf, ef)
g.nf = nf
g.ef = ef

return g
function FeatureGraph(fg::FeatureGraph; nf = fg.nf, ef = fg.ef,
senders = fg.senders, receivers = fg.receivers)
return FeatureGraph(nf, ef, senders, receivers)
end

"""
Expand Down Expand Up @@ -76,5 +61,6 @@ Aggregates the node features based on the given [`FeatureGraph`](@ref) and updat
"""
@inline function aggregate_node_features(graph::FeatureGraph, updated_edge_features)
return vcat(graph.nf,
NNlib.scatter(+, updated_edge_features, graph.receivers; dstsize = size(graph.nf)))
NNlib.scatter(+, updated_edge_features, graph.receivers;
dstsize = size(graph.nf), init = zero(eltype(graph.nf))))
end
98 changes: 21 additions & 77 deletions src/graph_net_blocks.jl
Original file line number Diff line number Diff line change
Expand Up @@ -3,91 +3,35 @@
# Licensed under the MIT license. See LICENSE file in the project root for details.
#

struct Encoder{T <: NamedTuple, N <: Lux.NAME_TYPE} <:
Lux.AbstractExplicitContainerLayer{(:layers,)}
layers::T
name::N
struct Encoder{N, E} <: Lux.AbstractLuxContainerLayer{(:node_layer, :edge_layer)}
node_layer::N
edge_layer::E
end

function Encoder(node_model, edge_model; name::Lux.NAME_TYPE = nothing)
fields = (Symbol("node_model_fn"), Symbol("edge_model_fn"))

return Encoder(NamedTuple{fields}((node_model, edge_model)), name)
end

function (e::Encoder)(graph::FeatureGraph, ps, st::NamedTuple{fields}) where {fields}
encode!(e.layers, graph, ps, st)
end

function encode!(
layers::NamedTuple{fields}, graph, ps, st::NamedTuple{fields}) where {fields}
nf, stn = layers[:node_model_fn](graph.nf, ps[:node_model_fn], st[:node_model_fn])
ef, ste = layers[:edge_model_fn](graph.ef, ps[:edge_model_fn], st[:edge_model_fn])
new_st = NamedTuple{fields}((stn, ste))

return update_features!(graph; nf = nf, ef = ef), new_st
end

struct Processor{T <: NamedTuple, N <: Lux.NAME_TYPE} <:
Lux.AbstractExplicitContainerLayer{(:layers,)}
layers::T
name::N
end

function Processor(node_model, edge_model; name::Lux.NAME_TYPE = nothing)
fields = (Symbol("node_model_fn"), Symbol("edge_model_fn"))

return Processor(NamedTuple{fields}((node_model, edge_model)), name)
end

function (p::Processor)(graph::FeatureGraph, ps, st::NamedTuple{fields}) where {fields}
process!(p.layers, graph, ps, st)
end

function process!(layers::NamedTuple{fields}, graph::FeatureGraph,
ps, st::NamedTuple{fields}) where {fields}
uef, ste = update_edge_features(
layers[:edge_model_fn], ps[:edge_model_fn], st[:edge_model_fn], graph)
unf, stn = update_node_features(
layers[:node_model_fn], ps[:node_model_fn], st[:node_model_fn], graph, uef)
new_st = NamedTuple{fields}((stn, ste))

return update_features!(graph; nf = graph.nf + unf, ef = graph.ef + uef), new_st
end

@inline function update_edge_features(el, ps, st, graph::FeatureGraph)
features = aggregate_edge_features(graph)

return el(features, ps, st)
end

@inline function update_node_features(
nl, ps, st, graph::FeatureGraph, updated_edge_features)
features = aggregate_node_features(graph, updated_edge_features)

return nl(features, ps, st)
function (e::Encoder)(graph::FeatureGraph, ps, st)
nf, stn = e.node_layer(graph.nf, ps.node_layer, st.node_layer)
ef, ste = e.edge_layer(graph.ef, ps.edge_layer, st.edge_layer)
return FeatureGraph(graph; nf = nf, ef = ef), (; node_layer = stn, edge_layer = ste)
end

struct Decoder{T <: NamedTuple, N <: Lux.NAME_TYPE} <:
Lux.AbstractExplicitContainerLayer{(:layers,)}
layers::T
name::N
struct Processor{N, E} <: Lux.AbstractLuxContainerLayer{(:node_layer, :edge_layer)}
node_layer::N
edge_layer::E
end

function Decoder(model; name::Lux.NAME_TYPE = nothing)
fields = (Symbol("model"),)

return Decoder(NamedTuple{fields}((model,)), name)
function (p::Processor)(graph::FeatureGraph, ps, st)
uef, ste = p.edge_layer(aggregate_edge_features(graph), ps.edge_layer, st.edge_layer)
unf, stn = p.node_layer(
aggregate_node_features(graph, uef), ps.node_layer, st.node_layer)
return FeatureGraph(graph; nf = graph.nf + unf, ef = graph.ef + uef),
(; node_layer = ste, edge_layer = stn)
end

function (d::Decoder)(graph::FeatureGraph, ps, st::NamedTuple{fields}) where {fields}
decode!(d.layers, graph, ps, st)
struct Decoder{D} <: Lux.AbstractLuxWrapperLayer{:decode_layer}
decode_layer::D
end

function decode!(layers::NamedTuple{fields}, graph::FeatureGraph,
ps, st::NamedTuple{fields}) where {fields}
y, stm = layers[:model](graph.nf, ps[:model], st[:model])
new_st = NamedTuple{fields}((stm,))

return y, new_st
function (d::Decoder)(graph::FeatureGraph, ps, st)
df, std = d.decode_layer(graph.nf, ps, st)
return df, std
end
Loading
Loading