Current section

Files

Jump to
viva_tensor src viva_tensor@nn@transformer.erl
Raw

src/viva_tensor@nn@transformer.erl

-module(viva_tensor@nn@transformer).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/nn/transformer.gleam").
-export([feed_forward_init/3, feed_forward_forward/2, encoder_block_init/4, encoder_block_forward/3, decoder_block_init/4, decoder_block_forward/3, transformer_init/6, transformer_encode/2, transformer_decode/3, transformer_forward/3]).
-export_type([activation/0, feed_forward/0, encoder_block/0, decoder_block/0, transformer/0]).
-if(?OTP_RELEASE >= 27).
-define(MODULEDOC(Str), -moduledoc(Str)).
-define(DOC(Str), -doc(Str)).
-else.
-define(MODULEDOC(Str), -compile([])).
-define(DOC(Str), -compile([])).
-endif.
?MODULEDOC(false).
-type activation() :: relu_act | gelu_act.
-type feed_forward() :: {feed_forward,
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
activation()}.
-type encoder_block() :: {encoder_block,
viva_tensor@nn@attention:multi_head_attention(),
feed_forward(),
viva_tensor@nn@norm:layer_norm(),
viva_tensor@nn@norm:layer_norm()}.
-type decoder_block() :: {decoder_block,
viva_tensor@nn@attention:multi_head_attention(),
viva_tensor@nn@attention:multi_head_attention(),
feed_forward(),
viva_tensor@nn@norm:layer_norm(),
viva_tensor@nn@norm:layer_norm(),
viva_tensor@nn@norm:layer_norm()}.
-type transformer() :: {transformer,
list(encoder_block()),
list(decoder_block()),
integer(),
integer()}.
-file("src/viva_tensor/nn/transformer.gleam", 67).
?DOC(false).
-spec apply_activation(viva_tensor@tensor:tensor(), activation()) -> viva_tensor@tensor:tensor().
apply_activation(T, Act) ->
case Act of
relu_act ->
viva_tensor@nn@activations:relu(T);
gelu_act ->
viva_tensor@nn@activations:gelu(T)
end.
-file("src/viva_tensor/nn/transformer.gleam", 103).
?DOC(false).
-spec feed_forward_init(integer(), integer(), activation()) -> feed_forward().
feed_forward_init(Embed_dim, Hidden_dim, Activation) ->
{feed_forward,
viva_tensor@tensor:zeros([Embed_dim, Hidden_dim]),
viva_tensor@tensor:zeros([Hidden_dim]),
viva_tensor@tensor:zeros([Hidden_dim, Embed_dim]),
viva_tensor@tensor:zeros([Embed_dim]),
Activation}.
-file("src/viva_tensor/nn/transformer.gleam", 139).
?DOC(false).
-spec add_bias_row(viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
add_bias_row(T, Bias) ->
case {erlang:element(3, T), erlang:element(3, Bias)} of
{[_, Out], [Bo]} when Bo =:= Out ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(T),
fun(T_data) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Bias),
fun(B_data) ->
New_data = begin
_pipe = gleam@list:sized_chunk(T_data, Out),
gleam@list:flat_map(
_pipe,
fun(Row) ->
gleam@list:map2(
Row,
B_data,
fun(X, B) -> X + B end
)
end
)
end,
{ok, {tensor, New_data, erlang:element(3, T)}}
end
)
end
);
{_, _} ->
{error,
{shape_mismatch, erlang:element(3, T), erlang:element(3, Bias)}}
end.
-file("src/viva_tensor/nn/transformer.gleam", 127).
?DOC(false).
-spec feed_forward_forward(feed_forward(), viva_tensor@tensor:tensor()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
feed_forward_forward(Ff, Input) ->
gleam@result:'try'(
viva_tensor@tensor:matmul(Input, erlang:element(2, Ff)),
fun(H1) ->
gleam@result:'try'(
add_bias_row(H1, erlang:element(3, Ff)),
fun(H1_b) ->
Activated = apply_activation(H1_b, erlang:element(6, Ff)),
gleam@result:'try'(
viva_tensor@tensor:matmul(
Activated,
erlang:element(4, Ff)
),
fun(H2) -> add_bias_row(H2, erlang:element(5, Ff)) end
)
end
)
end
).
-file("src/viva_tensor/nn/transformer.gleam", 183).
?DOC(false).
-spec encoder_block_init(integer(), integer(), integer(), activation()) -> {ok,
encoder_block()} |
{error, viva_tensor@core@error:tensor_error()}.
encoder_block_init(Embed_dim, Num_heads, Ffn_hidden_dim, Activation) ->
gleam@result:'try'(
viva_tensor@nn@attention:multi_head_attention_init(
Num_heads,
Embed_dim,
false
),
fun(Mha) ->
Ffn = feed_forward_init(Embed_dim, Ffn_hidden_dim, Activation),
{ok,
{encoder_block,
Mha,
Ffn,
viva_tensor@nn@norm:layer_norm_init(Embed_dim),
viva_tensor@nn@norm:layer_norm_init(Embed_dim)}}
end
).
-file("src/viva_tensor/nn/transformer.gleam", 211).
?DOC(false).
-spec encoder_block_forward(
encoder_block(),
viva_tensor@tensor:tensor(),
boolean()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
encoder_block_forward(Block, Input, Is_causal) ->
gleam@result:'try'(
viva_tensor@nn@norm:layer_norm_forward(erlang:element(4, Block), Input),
fun(X1) ->
gleam@result:'try'(
viva_tensor@nn@attention:multi_head_attention_forward(
erlang:element(2, Block),
X1,
X1,
X1,
Is_causal
),
fun(Attn_out) ->
gleam@result:'try'(
viva_tensor@tensor:add(Input, Attn_out),
fun(R1) ->
gleam@result:'try'(
viva_tensor@nn@norm:layer_norm_forward(
erlang:element(5, Block),
R1
),
fun(X2) ->
gleam@result:'try'(
feed_forward_forward(
erlang:element(3, Block),
X2
),
fun(Ffn_out) ->
viva_tensor@tensor:add(R1, Ffn_out)
end
)
end
)
end
)
end
)
end
).
-file("src/viva_tensor/nn/transformer.gleam", 265).
?DOC(false).
-spec decoder_block_init(integer(), integer(), integer(), activation()) -> {ok,
decoder_block()} |
{error, viva_tensor@core@error:tensor_error()}.
decoder_block_init(Embed_dim, Num_heads, Ffn_hidden_dim, Activation) ->
gleam@result:'try'(
viva_tensor@nn@attention:multi_head_attention_init(
Num_heads,
Embed_dim,
false
),
fun(Self_mha) ->
gleam@result:'try'(
viva_tensor@nn@attention:multi_head_attention_init(
Num_heads,
Embed_dim,
false
),
fun(Cross_mha) ->
Ffn = feed_forward_init(
Embed_dim,
Ffn_hidden_dim,
Activation
),
{ok,
{decoder_block,
Self_mha,
Cross_mha,
Ffn,
viva_tensor@nn@norm:layer_norm_init(Embed_dim),
viva_tensor@nn@norm:layer_norm_init(Embed_dim),
viva_tensor@nn@norm:layer_norm_init(Embed_dim)}}
end
)
end
).
-file("src/viva_tensor/nn/transformer.gleam", 596).
?DOC(false).
-spec range_loop(integer(), integer(), list(integer())) -> list(integer()).
range_loop(From, To, Acc) ->
case From > To of
true ->
lists:reverse(Acc);
false ->
range_loop(From + 1, To, [From | Acc])
end.
-file("src/viva_tensor/nn/transformer.gleam", 592).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/nn/transformer.gleam", 454).
?DOC(false).
-spec concat_heads(
list(viva_tensor@tensor:tensor()),
integer(),
integer(),
integer()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
concat_heads(Heads, Seq, Num_heads, Head_dim) ->
gleam@result:'try'(
gleam@list:try_map(
Heads,
fun(H) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(H),
fun(D) -> {ok, gleam@list:sized_chunk(D, Head_dim)} end
)
end
),
fun(Head_rows) ->
Combined = begin
_pipe = range_int(0, Seq - 1),
gleam@list:flat_map(_pipe, fun(S) -> _pipe@1 = Head_rows,
gleam@list:flat_map(
_pipe@1,
fun(Rows) -> _pipe@2 = Rows,
_pipe@3 = gleam@list:drop(_pipe@2, S),
_pipe@4 = gleam@list:first(_pipe@3),
gleam@result:unwrap(_pipe@4, []) end
) end)
end,
{ok, {tensor, Combined, [Seq, Num_heads * Head_dim]}}
end
).
-file("src/viva_tensor/nn/transformer.gleam", 430).
?DOC(false).
-spec split_heads(viva_tensor@tensor:tensor(), integer(), integer(), integer()) -> {ok,
list(viva_tensor@tensor:tensor())} |
{error, viva_tensor@core@error:tensor_error()}.
split_heads(X, Seq, Num_heads, Head_dim) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(X),
fun(Data) ->
Rows = gleam@list:sized_chunk(Data, Num_heads * Head_dim),
Heads = begin
_pipe = range_int(0, Num_heads - 1),
gleam@list:map(
_pipe,
fun(H) ->
Head_data = begin
_pipe@1 = Rows,
gleam@list:flat_map(
_pipe@1,
fun(Row) -> _pipe@2 = Row,
_pipe@3 = gleam@list:drop(
_pipe@2,
H * Head_dim
),
gleam@list:take(_pipe@3, Head_dim) end
)
end,
{tensor, Head_data, [Seq, Head_dim]}
end
)
end,
{ok, Heads}
end
).
-file("src/viva_tensor/nn/transformer.gleam", 414).
?DOC(false).
-spec rank2_dims(viva_tensor@tensor:tensor(), integer(), binary()) -> {ok,
{integer(), integer()}} |
{error, viva_tensor@core@error:tensor_error()}.
rank2_dims(T, Embed_dim, Who) ->
case erlang:element(3, T) of
[Seq, E] when E =:= Embed_dim ->
{ok, {Seq, E}};
_ ->
_ = Who,
{error, {shape_mismatch, [-1, Embed_dim], erlang:element(3, T)}}
end.
-file("src/viva_tensor/nn/transformer.gleam", 351).
?DOC(false).
-spec cross_attention_forward(
viva_tensor@nn@attention:multi_head_attention(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
cross_attention_forward(Mha, Q, K, V) ->
gleam@result:'try'(
rank2_dims(Q, erlang:element(3, Mha), <<"cross_attn q"/utf8>>),
fun(_use0) ->
{Seq_q, _} = _use0,
gleam@result:'try'(
rank2_dims(K, erlang:element(3, Mha), <<"cross_attn k"/utf8>>),
fun(_use0@1) ->
{Seq_k, _} = _use0@1,
gleam@result:'try'(
rank2_dims(
V,
erlang:element(3, Mha),
<<"cross_attn v"/utf8>>
),
fun(_use0@2) ->
{Seq_v, _} = _use0@2,
case Seq_k =:= Seq_v of
false ->
{error,
{shape_mismatch,
erlang:element(3, K),
erlang:element(3, V)}};
true ->
gleam@result:'try'(
viva_tensor@tensor:matmul(
Q,
erlang:element(5, Mha)
),
fun(Q_proj) ->
gleam@result:'try'(
viva_tensor@tensor:matmul(
K,
erlang:element(6, Mha)
),
fun(K_proj) ->
gleam@result:'try'(
viva_tensor@tensor:matmul(
V,
erlang:element(
7,
Mha
)
),
fun(V_proj) ->
gleam@result:'try'(
split_heads(
Q_proj,
Seq_q,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(Q_heads) ->
gleam@result:'try'(
split_heads(
K_proj,
Seq_k,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(
K_heads
) ->
gleam@result:'try'(
split_heads(
V_proj,
Seq_v,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(
V_heads
) ->
Triples = begin
_pipe = gleam@list:zip(
Q_heads,
gleam@list:zip(
K_heads,
V_heads
)
),
gleam@list:map(
_pipe,
fun(
Triple
) ->
{Qh,
Rest} = Triple,
{Kh,
Vh} = Rest,
{Qh,
Kh,
Vh}
end
)
end,
gleam@result:'try'(
gleam@list:try_map(
Triples,
fun(
T
) ->
{Qh@1,
Kh@1,
Vh@1} = T,
viva_tensor@nn@attention:scaled_dot_product_attention(
Qh@1,
Kh@1,
Vh@1,
none,
false
)
end
),
fun(
Head_outputs
) ->
gleam@result:'try'(
concat_heads(
Head_outputs,
Seq_q,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(
Concat
) ->
viva_tensor@tensor:matmul(
Concat,
erlang:element(
8,
Mha
)
)
end
)
end
)
end
)
end
)
end
)
end
)
end
)
end
)
end
end
)
end
)
end
).
-file("src/viva_tensor/nn/transformer.gleam", 316).
?DOC(false).
-spec decoder_block_forward(
decoder_block(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
decoder_block_forward(Block, Input, Encoder_output) ->
gleam@result:'try'(
viva_tensor@nn@norm:layer_norm_forward(erlang:element(5, Block), Input),
fun(X1) ->
gleam@result:'try'(
viva_tensor@nn@attention:multi_head_attention_forward(
erlang:element(2, Block),
X1,
X1,
X1,
true
),
fun(Self_out) ->
gleam@result:'try'(
viva_tensor@tensor:add(Input, Self_out),
fun(R1) ->
gleam@result:'try'(
viva_tensor@nn@norm:layer_norm_forward(
erlang:element(6, Block),
R1
),
fun(X2) ->
gleam@result:'try'(
cross_attention_forward(
erlang:element(3, Block),
X2,
Encoder_output,
Encoder_output
),
fun(Cross_out) ->
gleam@result:'try'(
viva_tensor@tensor:add(
R1,
Cross_out
),
fun(R2) ->
gleam@result:'try'(
viva_tensor@nn@norm:layer_norm_forward(
erlang:element(
7,
Block
),
R2
),
fun(X3) ->
gleam@result:'try'(
feed_forward_forward(
erlang:element(
4,
Block
),
X3
),
fun(Ffn_out) ->
viva_tensor@tensor:add(
R2,
Ffn_out
)
end
)
end
)
end
)
end
)
end
)
end
)
end
)
end
).
-file("src/viva_tensor/nn/transformer.gleam", 504).
?DOC(false).
-spec transformer_init(
integer(),
integer(),
integer(),
integer(),
integer(),
activation()
) -> {ok, transformer()} | {error, viva_tensor@core@error:tensor_error()}.
transformer_init(
Num_encoder_layers,
Num_decoder_layers,
Embed_dim,
Num_heads,
Ffn_hidden_dim,
Activation
) ->
case (Num_encoder_layers < 0) orelse (Num_decoder_layers < 0) of
true ->
{error,
{invalid_shape,
<<<<<<<<"transformer_init: layer counts must be non-negative (got num_encoder_layers="/utf8,
(erlang:integer_to_binary(
Num_encoder_layers
))/binary>>/binary,
", num_decoder_layers="/utf8>>/binary,
(erlang:integer_to_binary(Num_decoder_layers))/binary>>/binary,
")"/utf8>>}};
false ->
gleam@result:'try'(
begin
_pipe = gleam@list:repeat(nil, Num_encoder_layers),
gleam@list:try_map(
_pipe,
fun(_) ->
encoder_block_init(
Embed_dim,
Num_heads,
Ffn_hidden_dim,
Activation
)
end
)
end,
fun(Encoders) ->
gleam@result:'try'(
begin
_pipe@1 = gleam@list:repeat(nil, Num_decoder_layers),
gleam@list:try_map(
_pipe@1,
fun(_) ->
decoder_block_init(
Embed_dim,
Num_heads,
Ffn_hidden_dim,
Activation
)
end
)
end,
fun(Decoders) ->
{ok,
{transformer,
Encoders,
Decoders,
Num_encoder_layers,
Num_decoder_layers}}
end
)
end
)
end.
-file("src/viva_tensor/nn/transformer.gleam", 550).
?DOC(false).
-spec transformer_encode(transformer(), viva_tensor@tensor:tensor()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
transformer_encode(Model, Src) ->
gleam@list:try_fold(
erlang:element(2, Model),
Src,
fun(Acc, Block) -> encoder_block_forward(Block, Acc, false) end
).
-file("src/viva_tensor/nn/transformer.gleam", 566).
?DOC(false).
-spec transformer_decode(
transformer(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
transformer_decode(Model, Tgt, Memory) ->
gleam@list:try_fold(
erlang:element(3, Model),
Tgt,
fun(Acc, Block) -> decoder_block_forward(Block, Acc, Memory) end
).
-file("src/viva_tensor/nn/transformer.gleam", 583).
?DOC(false).
-spec transformer_forward(
transformer(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
transformer_forward(Model, Src, Tgt) ->
gleam@result:'try'(
transformer_encode(Model, Src),
fun(Memory) -> transformer_decode(Model, Tgt, Memory) end
).