Skip to content
Open
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
3 changes: 3 additions & 0 deletions src/LuxRecurrentLayers.jl
Original file line number Diff line number Diff line change
Expand Up @@ -19,12 +19,15 @@ BoolType = Utils.BoolType

@compat(public, (initialparameters, initialstates, parameterlength, statelength))

export AdditiveIntegration, MultiplicativeIntegration

export AntisymmetricRNNCell, ATRCell, BRCell, CFNCell, coRNNCell, FastGRNNCell,
FastRNNCell, GatedAntisymmetricRNNCell, IndRNNCell, JANETCell, LEMCell, LightRUCell,
LiGRUCell, MGUCell, MinimalRNNCell, MultiplicativeLSTMCell, MUT1Cell, MUT2Cell,
MUT3Cell, NASCell, NBRCell, PeepholeLSTMCell, RANCell, SCRNCell, SGRNCell,
STARCell, TGRUCell, TLSTMCell, TRNNCell, UnICORNNCell, WMCLSTMCell

include("base_functions.jl")
include("generics.jl")

include("cells/antisymmetricrnn_cell.jl")
Expand Down
24 changes: 24 additions & 0 deletions src/base_functions.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
struct AdditiveIntegration end
struct MultiplicativeIntegration end

function recurrence_double_bias(integration_mode, wih::AbstractMatrix, whh::AbstractMatrix,
inp::AbstractArray, state::AbstractArray, bih::Union{AbstractVector, Nothing},
bhh::Union{AbstractVector, Nothing}, bmi::Union{AbstractVector, Nothing};
activation_inp = identity, activation_state = identity, activation_recurrence = tanh_fast)
wih_inp_bih = fused_dense_bias_activation(activation_inp, wih, inp, bih)
whh_state_bhh = fused_dense_bias_activation(activation_state, whh, state, bhh)

return dense_integration(integration_mode, wih_inp_bih, whh_state_bhh, bmi; activation = activation_recurrence)
end

function dense_integration(::AdditiveIntegration, wih_inp_bih::AbstractMatrix,
whh_state_bhh::AbstractMatrix, bmi::Union{AbstractVector, Nothing}; activation=identity)

return bias_activation(activation, wih_inp_bih .+ whh_state_bhh, bmi)
end

function dense_integration(::MultiplicativeIntegration, wih_inp_bih::AbstractMatrix,
whh_state_bhh::AbstractMatrix, bmi::Union{AbstractVector, Nothing}; activation=identity)

return bias_activation(activation, wih_inp_bih .* whh_state_bhh, bmi)
end
73 changes: 49 additions & 24 deletions src/cells/antisymmetricrnn_cell.jl
Original file line number Diff line number Diff line change
@@ -1,11 +1,11 @@
#https://arxiv.org/abs/1902.09689
@doc raw"""
AntisymmetricRNNCell(in_dims => out_dims, [activation];
use_bias=true, use_recurrent_bias=true,
train_state=false, init_bias=nothing,
init_recurrent_bias=nothing, init_weight=nothing,
use_bias=true, use_recurrent_bias=true, use_integration_bias=false,
train_state=false, init_bias=nothing, init_recurrent_bias=nothing,
init_integration_bias=nothing, init_weight=nothing,
init_recurrent_weight=nothing, init_state=zeros32,
epsilon=1.0, gamma=0.0)
epsilon=1.0, gamma=0.0, integration_mode=AdditiveIntegration())


[Antisymmetric recurrent cell](https://arxiv.org/abs/1902.09689).
Expand Down Expand Up @@ -33,14 +33,18 @@
Default set to `true`.
- `use_recurrent_bias`: Flag to use recurrent bias $\mathbf{b}_{hh}$ in the computation.
Default set to `true`.
- `use_recurrent_bias`: Flag to use recurrent bias $\mathbf{b}_{hh}$ in the computation.
Default set to `true`.
- `use_integration_bias`: Flag to use integration bias $\mathbf{b}_{mi}$ in the computation.
This bias is only useful for multiplicative integration. Check the docs page on multiplicative
integration for more details. Default set to `false`.
- `train_state`: Flag to set the initial hidden state as trainable.
Default set to `false`.
- `init_bias`: Initializer for bias $\mathbf{b}_{ih}$. If set to
`nothing`, weights are initialized from a uniform distribution within `[-bound, bound]`
where `bound = inv(sqrt(out_dims))`. Default is `nothing`.
- `init_bias`: Initializer for recurrent bias $\mathbf{b}_{hh}$. If set to
- `init_recurrent_bias`: Initializer for recurrent bias $\mathbf{b}_{hh}$. If set to
`nothing`, weights are initialized from a uniform distribution within `[-bound, bound]`
where `bound = inv(sqrt(out_dims))`. Default is `nothing`.
- `init_integration_bias`: Initializer for integration bias $\mathbf{b}_{mi}$. If set to
`nothing`, weights are initialized from a uniform distribution within `[-bound, bound]`
where `bound = inv(sqrt(out_dims))`. Default is `nothing`.
- `init_weight`: Initializer for weight $\mathbf{W}_{ih}$. If set to
Expand All @@ -52,6 +56,8 @@
- `init_state`: Initializer for hidden state. Default set to `zeros32`.
- `epsilon`: step size $\epsilon$. Default is 1.0.
- `gamma`: strength of diffusion $\gamma$. Default is 0.0.
- `integration_mode`: integration type for the recurrent forward pass.
Default is [`AdditiveIntegration()`](@ref).

## Inputs

Expand Down Expand Up @@ -80,7 +86,9 @@
- `bias_ih`: Bias vector for the input-hidden connection (not present if
`use_bias=false`) $\mathbf{b}_{ih}$.
- `bias_hh`: Bias vector for the hidden-hidden connection (not present if
`use_bias=false`) $\mathbf{b}_{hh}$.
`use_recurrent_bias=false`) $\mathbf{b}_{hh}$.
- `bias_mi`: Bias vector for the integration connection (not present if `use_integration_bias=false`)
$\mathbf{b}_{mi}$
- `hidden_state`: Initial hidden state vector (not present if `train_state=false`)

## States
Expand All @@ -95,24 +103,26 @@
out_dims <: IntegerType
init_bias
init_recurrent_bias
init_integration_bias
init_weight
init_recurrent_weight
init_state
use_bias <: StaticBool
use_recurrent_bias <: StaticBool
epsilon
gamma
integration_mode
end

function AntisymmetricRNNCell(
(in_dims, out_dims)::Pair{<:IntegerType, <:IntegerType}, activation=tanh;
use_bias::BoolType=True(), use_recurrent_bias::BoolType=True(),
use_bias::BoolType=True(), use_recurrent_bias::BoolType=True(), use_integration_bias::BoolType=False(),
train_state::BoolType=False(), init_bias=nothing, init_recurrent_bias=nothing,
init_weight=nothing, init_recurrent_weight=nothing, init_state=zeros32,
epsilon=1.0f0, gamma=0.0f0)
init_integration_bias=nothing, init_weight=nothing, init_recurrent_weight=nothing, init_state=zeros32,
epsilon=1.0f0, gamma=0.0f0, integration_mode=AdditiveIntegration())
return AntisymmetricRNNCell(static(train_state), activation, in_dims, out_dims,
init_bias, init_recurrent_bias, init_weight, init_recurrent_weight, init_state,
static(use_bias), static(use_recurrent_bias), epsilon, gamma)
init_bias, init_recurrent_bias, init_integration_bias, init_weight, init_recurrent_weight, init_state,
static(use_bias), static(use_recurrent_bias), epsilon, gamma, integration_mode)
end

function initialparameters(rng::AbstractRNG, asymrnn::AntisymmetricRNNCell)
Expand All @@ -129,12 +139,12 @@ function (asymrnn::AntisymmetricRNNCell)(
ps, st::NamedTuple)
matched_inp, matched_state = match_eltype(asymrnn, ps, st, inp, state)
bias_ih = safe_getproperty(ps, Val(:bias_ih))
linear_input = fused_dense_bias_activation(identity, ps.weight_ih, matched_inp, bias_ih)
bias_hh = safe_getproperty(ps, Val(:bias_hh))
bias_mi = safe_getproperty(ps, Val(:bias_mi))
asym_weight_hh = compute_asym_recurrent(ps.weight_hh, asymrnn.gamma)
linear_recur = fused_dense_bias_activation(
identity, asym_weight_hh, matched_state, bias_hh)
half_new_state = fast_activation!!(asymrnn.activation, linear_input .+ linear_recur)
full_gs = recurrence_double_bias(ligru.integration_mode, ps.weight_ih, asym_weight_hh,
matched_inp, matched_state, bias_ih, bias_hh, bias_mi)
half_new_state = fast_activation!!(asymrnn.activation, full_gs)
new_state = matched_state .+ asymrnn.epsilon .* half_new_state
return (new_state, (new_state,)), st
end
Expand All @@ -143,17 +153,19 @@ function Base.show(io::IO, r::AntisymmetricRNNCell)
print(io, "AntisymmetricRNNCell($(r.in_dims) => $(r.out_dims)")
(r.activation == identity) || print(io, ", $(r.activation)")
has_bias(r) || print(io, ", use_bias=false")
has_recurrent_bias(r) || print(io, ", use_recurrent_bias=false")
has_integration_bias(r) || print(io, ", use_integration_bias=false")
has_train_state(r) && print(io, ", train_state=true")
print(io, ")")
end

@doc raw"""
GatedAntisymmetricRNNCell(in_dims => out_dims, [activation];
use_bias=true, use_recurrent_bias=true,
train_state=false, init_bias=nothing,
init_recurrent_bias=nothing, init_weight=nothing,
use_bias=true, use_recurrent_bias=true, use_integration_bias=false,
train_state=false, init_bias=nothing, init_recurrent_bias=nothing,
init_integration_bias=nothing, init_weight=nothing,
init_recurrent_weight=nothing, init_state=zeros32,
epsilon=1.0, gamma=0.0)
epsilon=1.0, gamma=0.0, integration_mode=AdditiveIntegration())



Expand Down Expand Up @@ -188,6 +200,9 @@ end
Default set to `true`.
- `use_recurrent_bias`: Flag to use recurrent bias $\mathbf{b}_{hh}$ in the computation.
Default set to `true`.
- `use_integration_bias`: Flag to use integration bias $\mathbf{b}_{mi}$ in the computation.
This bias is only useful for multiplicative integration. Check the docs page on multiplicative
integration for more details. Default set to `false`.
- `train_state`: Flag to set the initial hidden state as trainable.
Default set to `false`.
- `init_bias`: Initializer for input to hidden bias $\mathbf{b}_{ih}^z, \mathbf{b}_{ih}^h$.
Expand All @@ -199,6 +214,9 @@ end
- `init_recurrent_bias`: Initializer for hidden to hidden bias $\mathbf{b}_{hh}$. If set to `nothing`,
weights are initialized from a uniform distribution within `[-bound, bound]` where
`bound = inv(sqrt(out_dims))`. Default is `nothing`.
- `init_integration_bias`: Initializer for integration bias $\mathbf{b}_{mi}$. If set to
`nothing`, weights are initialized from a uniform distribution within `[-bound, bound]`
where `bound = inv(sqrt(out_dims))`. Default is `nothing`.
- `init_weight`: Initializer for input to hidden weights $\mathbf{W}_{ih}^z, \mathbf{W}_{ih}^x$.
Must be a tuple containing 2 functions, e.g., `(glorot_normal, kaiming_uniform)`.
If a single function `fn` is provided, it is automatically expanded into
Expand All @@ -211,6 +229,8 @@ end
- `init_state`: Initializer for hidden state. Default set to `zeros32`.
- `epsilon`: step size. Default is 1.0.
- `gamma`: strength of diffusion. Default is 0.0.
- `integration_mode`: integration type for the recurrent forward pass.
Default is [`AdditiveIntegration()`](@ref).

## Inputs

Expand Down Expand Up @@ -243,8 +263,10 @@ end
``\{ \mathbf{b}_{ih}^z, \mathbf{b}_{ih}^h \}``
The initializers in `init_bias` are applied in the order they appear:
the first function is used for $\mathbf{b}_{ih}^z$, and the second for $\mathbf{b}_{ih}^h$.
- `bias_hh`: Bias vector for the hidden-hidden connection (not present if `use_bias=false`)
- `bias_hh`: Bias vector for the hidden-hidden connection (not present if `use_recurrent_bias=false`)
$\mathbf{b}_{hh}$
- `bias_mi`: Bias vector for the integration connection (not present if `use_integration_bias=false`)
$\mathbf{b}_{mi}$
- `hidden_state`: Initial hidden state vector (not present if `train_state=false`)

## States
Expand Down Expand Up @@ -312,12 +334,13 @@ function (asymrnn::GatedAntisymmetricRNNCell)(
matched_inp, matched_state = match_eltype(asymrnn, ps, st, inp, state)
bias_ih = safe_getproperty(ps, Val(:bias_ih))
bias_hh = safe_getproperty(ps, Val(:bias_hh))
bias_mi = safe_getproperty(ps, Val(:bias_mi))
full_gxs = fused_dense_bias_activation(identity, ps.weight_ih, matched_inp, bias_ih)
gxs = multigate(full_gxs, Val(2))
asym_weight_hh = compute_asym_recurrent(ps.weight_hh, asymrnn.gamma)
hs = fused_dense_bias_activation(identity, asym_weight_hh, matched_state, bias_hh)
input_gate = @. sigmoid_fast(hs + gxs[1])
half_new_state = @. tanh_fast(hs + gxs[2])
input_gate = sigmoid_fast.(dense_integration(asymrnn.integration_mode, hs, gxs[1], bias_mi))
half_new_state = @. tanh_fast(dense_integration(asymrnn.integration_mode, hs, gxs[2], bias_mi))
new_state = @. matched_state .+ asymrnn.epsilon .* input_gate
return (new_state, (new_state,)), st
end
Expand All @@ -326,6 +349,8 @@ function Base.show(io::IO, r::GatedAntisymmetricRNNCell)
print(io, "GatedAntisymmetricRNNCell($(r.in_dims) => $(r.out_dims)")
(r.activation == identity) || print(io, ", $(r.activation)")
has_bias(r) || print(io, ", use_bias=false")
has_recurrent_bias(r) || print(io, ", use_recurrent_bias=false")
has_integration_bias(r) || print(io, ", use_integration_bias=false")
has_train_state(r) && print(io, ", train_state=true")
print(io, ")")
end
Expand Down
44 changes: 34 additions & 10 deletions src/cells/br_cell.jl
Original file line number Diff line number Diff line change
@@ -1,9 +1,10 @@
#https://doi.org/10.1371/journal.pone.0252676
@doc raw"""
BRCell(in_dims => out_dims;
use_bias=true, use_recurrent_bias=true, train_state=false, init_bias=nothing,
use_bias=true, use_recurrent_bias=true, use_integration_bias=false,
train_state=false, init_bias=nothing, init_recurrent_bias=nothing,
init_weight=nothing, init_recurrent_weight=nothing,
init_state=zeros32)
init_state=zeros32, integration_mode=AdditiveIntegration())

[Bistable recurrent cell](https://doi.org/10.1371/journal.pone.0252676).

Expand Down Expand Up @@ -34,6 +35,9 @@
Default set to `true`.
- `use_recurrent_bias`: Flag to use recurrent bias $\mathbf{b}_{hh}$ in the computation.
Default set to `true`.
- `use_integration_bias`: Flag to use integration bias $\mathbf{b}_{mi}$ in the computation.
This bias is only useful for multiplicative integration. Check the docs page on multiplicative
integration for more details. Default set to `false`.
- `train_state`: Flag to set the initial hidden state as trainable.
Default set to `false`.
- `init_bias`: Initializer for input to hidden bias
Expand All @@ -50,6 +54,12 @@
2-element tuple (fn, fn). If set to `nothing`, weights are initialized from a
uniform distribution within `[-bound, bound]` where `bound = inv(sqrt(out_dims))`.
Default is `nothing`.
- `init_integration_bias`: Initializer for integration bias $\mathbf{b}_{mi}$.
Must be a tuple containing 2 functions, e.g., `(glorot_normal, kaiming_uniform)`.
If a single function `fn` is provided, it is automatically expanded into a
2-element tuple (fn, fn). If set to `nothing`, weights are initialized from a
uniform distribution within `[-bound, bound]` where `bound = inv(sqrt(out_dims))`.
Default is `nothing`.
- `init_weight`: Initializer for input to hidden weights
$\mathbf{W}_{ih}^a, \mathbf{W}_{ih}^c, \mathbf{W}_{ih}^h$.
Must be a tuple containing 3 functions, e.g., `(glorot_normal, kaiming_uniform)`.
Expand All @@ -65,6 +75,8 @@
a uniform distribution within `[-bound, bound]` where `bound = inv(sqrt(out_dims))`.
Default is `nothing`.
- `init_state`: Initializer for hidden state. Default set to `zeros32`.
- `integration_mode`: integration type for the recurrent forward pass.
Default is [`AdditiveIntegration()`](@ref).

## Inputs

Expand Down Expand Up @@ -108,6 +120,8 @@
The initializers in `init_bias` are applied in the order they appear:
the first function is used for $\mathbf{b}_{hh}^z$, and the second for
$\mathbf{b}_{hh}^c$.
- `bias_mi`: Bias vector for the integration connection (not present if `use_integration_bias=false`)
$\mathbf{b}_{mi}$
- `hidden_state`: Initial hidden state vector (not present if `train_state=false`)

## States
Expand All @@ -121,27 +135,31 @@
out_dims <: IntegerType
init_bias
init_recurrent_bias
init_integration_bias
init_weight
init_recurrent_weight
init_state
use_bias <: StaticBool
use_recurrent_bias <: StaticBool
integration_mode
end

function BRCell((in_dims, out_dims)::Pair{<:IntegerType, <:IntegerType};
use_bias::BoolType=True(), use_recurrent_bias::BoolType=True(),
use_bias::BoolType=True(), use_recurrent_bias::BoolType=True(), use_integration_bias::BoolType=False(),
train_state::BoolType=False(), init_bias=nothing,
init_recurrent_bias=nothing, init_weight=nothing, init_recurrent_weight=nothing,
init_state=zeros32)
init_recurrent_bias=nothing, init_integration_bias=nothing, init_weight=nothing, init_recurrent_weight=nothing,
init_state=zeros32, integration_mode=AdditiveIntegration())
init_weight isa NTuple{3} || (init_weight = ntuple(Returns(init_weight), 3))
init_recurrent_weight isa NTuple{2} ||
(init_recurrent_weight = ntuple(Returns(init_recurrent_weight), 2))
init_bias isa NTuple{3} || (init_bias = ntuple(Returns(init_bias), 3))
init_recurrent_bias isa NTuple{3} ||
(init_recurrent_bias = ntuple(Returns(init_recurrent_bias), 3))
init_integration_bias isa NTuple{2} ||
(init_integration_bias = ntuple(Returns(init_integration_bias), 2))
return BRCell(static(train_state), in_dims, out_dims, init_bias, init_recurrent_bias,
init_weight, init_recurrent_weight, init_state, static(use_bias),
static(use_recurrent_bias))
init_integration_bias, init_weight, init_recurrent_weight, init_state, static(use_bias),
static(use_recurrent_bias), integration_mode)
end

function initialparameters(rng::AbstractRNG, br::BRCell)
Expand All @@ -156,6 +174,9 @@ function initialparameters(rng::AbstractRNG, br::BRCell)
elseif has_recurrent_bias(br)
bias_hh = multi_bias(rng, br.init_recurrent_bias, br.out_dims, br.out_dims)
ps = merge(ps, (; bias_hh))
elseif has_integration_bias(br)
bias_mi = multi_bias(rng, br.init_integration_bias, br.out_dims, br.out_dims)
ps = merge(ps, (; bias_mi))
end
has_train_state(br) &&
(ps = merge(ps, (hidden_state=br.init_state(rng, br.out_dims),)))
Expand All @@ -175,14 +196,17 @@ function (br::BRCell)(
matched_inp, matched_state = match_eltype(br, ps, st, inp, state)
bias_ih = safe_getproperty(ps, Val(:bias_ih))
bias_hh = safe_getproperty(ps, Val(:bias_hh))
bias_mi = safe_getproperty(ps, Val(:bias_mi))
t_ones = one(eltype(matched_inp))
full_xs = fused_dense_bias_activation(identity, ps.weight_ih, matched_inp, bias_ih)
xs = multigate(full_xs, Val(3))
ws = multigate(ps.weight_hh, Val(2))
bhs = bias_safe_multigate(bias_hh, Val(3))
modulation_gate = t_ones .+
bias_activation(tanh_fast, xs[1] .+ ws[1] .* matched_state, bhs[1])
candidate_state = bias_activation(sigmoid_fast, xs[2] .+ ws[2] .* matched_state, bhs[2])
bmis = bias_safe_multigate(bias_mi, Val(2))
wh_state_1 = fused_dense_bias_activation(identity, ws[1], matched_state, bhs[1])
wh_state_2 = fused_dense_bias_activation(identity, ws[2], matched_state, bhs[2])
modulation_gate = t_ones .+ dense_integration(br.integration_mode, xs[1], wh_state_1, bmis[1])
candidate_state = dense_integration(br.integration_mode, xs[2], wh_state_1, bmis[2]; activation=sigmoid_fast)
new_state = candidate_state .* matched_state .+
(t_ones .- candidate_state) .*
bias_activation(
Expand Down
Loading
Loading