Current section

Files

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

src/viva_tensor@nn@attention.erl

-module(viva_tensor@nn@attention).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/nn/attention.gleam").
-export([causal_mask/1, scaled_dot_product_attention/5, multi_head_attention_init/3, multi_head_attention_forward/5]).
-export_type([multi_head_attention/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 multi_head_attention() :: {multi_head_attention,
integer(),
integer(),
integer(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
gleam@option:option(viva_tensor@tensor:tensor()),
gleam@option:option(viva_tensor@tensor:tensor()),
gleam@option:option(viva_tensor@tensor:tensor()),
gleam@option:option(viva_tensor@tensor:tensor())}.
-file("src/viva_tensor/nn/attention.gleam", 168).
?DOC(false).
-spec crop_2d(list(float()), integer(), integer(), integer(), integer()) -> list(float()).
crop_2d(Data, Rows, Cols, New_rows, New_cols) ->
_ = Rows,
_pipe = Data,
_pipe@1 = gleam@list:sized_chunk(_pipe, Cols),
_pipe@2 = gleam@list:take(_pipe@1, New_rows),
_pipe@3 = gleam@list:map(
_pipe@2,
fun(Row) -> gleam@list:take(Row, New_cols) end
),
lists:append(_pipe@3).
-file("src/viva_tensor/nn/attention.gleam", 134).
?DOC(false).
-spec apply_mask(
viva_tensor@tensor:tensor(),
gleam@option:option(viva_tensor@tensor:tensor()),
integer(),
integer()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
apply_mask(Scores, Mask, Seq_q, Seq_k) ->
case Mask of
none ->
{ok, Scores};
{some, M} ->
case erlang:element(3, M) of
[Mq, Mk] when (Mq >= Seq_q) andalso (Mk >= Seq_k) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Scores),
fun(Score_data) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(M),
fun(Mask_data) ->
Cropped_mask = case (Mq =:= Seq_q) andalso (Mk
=:= Seq_k) of
true ->
Mask_data;
false ->
crop_2d(
Mask_data,
Mq,
Mk,
Seq_q,
Seq_k
)
end,
Merged = gleam@list:map2(
Score_data,
Cropped_mask,
fun(S, Mv) -> case Mv > 0.5 of
true ->
S;
false ->
S + -1.0e9
end end
),
{ok, {tensor, Merged, [Seq_q, Seq_k]}}
end
)
end
);
_ ->
{error,
{shape_mismatch, [Seq_q, Seq_k], erlang:element(3, M)}}
end
end.
-file("src/viva_tensor/nn/attention.gleam", 476).
?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/attention.gleam", 472).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/nn/attention.gleam", 190).
?DOC(false).
-spec causal_mask(integer()) -> viva_tensor@tensor:tensor().
causal_mask(Seq_len) ->
Data = begin
_pipe = range_int(0, Seq_len - 1),
gleam@list:flat_map(
_pipe,
fun(I) -> _pipe@1 = range_int(0, Seq_len - 1),
gleam@list:map(_pipe@1, fun(J) -> case J =< I of
true ->
1.0;
false ->
+0.0
end end) end
)
end,
{tensor, Data, [Seq_len, Seq_len]}.
-file("src/viva_tensor/nn/attention.gleam", 93).
?DOC(false).
-spec sdpa_run(
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
gleam@option:option(viva_tensor@tensor:tensor()),
boolean(),
integer(),
integer(),
integer()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
sdpa_run(Q, K, V, Mask, Is_causal, Seq_q, Seq_k, Dim) ->
Scale = case gleam@float:square_root(erlang:float(Dim)) of
{ok, S} when S > +0.0 ->
case S of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> 1.0 / Gleam@denominator
end;
_ ->
1.0
end,
gleam@result:'try'(
viva_tensor@tensor:transpose(K),
fun(K_t) ->
gleam@result:'try'(
viva_tensor@tensor:matmul(Q, K_t),
fun(Raw_scores) ->
Scaled_scores = viva_tensor@tensor:scale(Raw_scores, Scale),
Effective_mask = case Is_causal of
true ->
{some, causal_mask(gleam@int:max(Seq_q, Seq_k))};
false ->
Mask
end,
gleam@result:'try'(
apply_mask(Scaled_scores, Effective_mask, Seq_q, Seq_k),
fun(Masked_scores) ->
gleam@result:'try'(
viva_tensor@tensor:softmax_axis(
Masked_scores,
1
),
fun(Weights) ->
viva_tensor@tensor:matmul(Weights, V)
end
)
end
)
end
)
end
).
-file("src/viva_tensor/nn/attention.gleam", 67).
?DOC(false).
-spec scaled_dot_product_attention(
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
gleam@option:option(viva_tensor@tensor:tensor()),
boolean()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
scaled_dot_product_attention(Q, K, V, Mask, Is_causal) ->
case {erlang:element(3, Q), erlang:element(3, K), erlang:element(3, V)} of
{[Seq_q, Dim_q], [Seq_k, Dim_k], [Seq_v, _]} ->
case Dim_q =:= Dim_k of
false ->
{error, {shape_mismatch, [Seq_q, Dim_q], [Seq_k, Dim_k]}};
true ->
case Seq_k =:= Seq_v of
false ->
{error,
{shape_mismatch,
erlang:element(3, K),
erlang:element(3, V)}};
true ->
sdpa_run(
Q,
K,
V,
Mask,
Is_causal,
Seq_q,
Seq_k,
Dim_q
)
end
end;
{_, _, _} ->
{error,
{invalid_shape,
<<"scaled_dot_product_attention expects rank-2 q/k/v tensors"/utf8>>}}
end.
-file("src/viva_tensor/nn/attention.gleam", 240).
?DOC(false).
-spec multi_head_attention_init(integer(), integer(), boolean()) -> {ok,
multi_head_attention()} |
{error, viva_tensor@core@error:tensor_error()}.
multi_head_attention_init(Num_heads, Embed_dim, Use_bias) ->
case (Num_heads =< 0) orelse (Embed_dim =< 0) of
true ->
{error,
{invalid_shape,
<<"multi_head_attention: num_heads and embed_dim must be positive"/utf8>>}};
false ->
case case Num_heads of
0 -> 0;
Gleam@denominator -> Embed_dim rem Gleam@denominator
end of
0 ->
Head_dim = case Num_heads of
0 -> 0;
Gleam@denominator@1 -> Embed_dim div Gleam@denominator@1
end,
W = viva_tensor@tensor:zeros([Embed_dim, Embed_dim]),
Bias = case Use_bias of
true ->
{some, viva_tensor@tensor:zeros([Embed_dim])};
false ->
none
end,
{ok,
{multi_head_attention,
Num_heads,
Embed_dim,
Head_dim,
W,
W,
W,
W,
Bias,
Bias,
Bias,
Bias}};
_ ->
{error,
{invalid_shape,
<<<<<<<<"multi_head_attention: embed_dim ("/utf8,
(erlang:integer_to_binary(Embed_dim))/binary>>/binary,
") not divisible by num_heads ("/utf8>>/binary,
(erlang:integer_to_binary(Num_heads))/binary>>/binary,
")"/utf8>>}}
end
end.
-file("src/viva_tensor/nn/attention.gleam", 396).
?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
{[Seq, 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 gleam@float:add/2
)
end
)
end,
_ = Seq,
{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/attention.gleam", 383).
?DOC(false).
-spec linear(
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
gleam@option:option(viva_tensor@tensor:tensor())
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
linear(X, W, B) ->
gleam@result:'try'(viva_tensor@tensor:matmul(X, W), fun(Y) -> case B of
none ->
{ok, Y};
{some, Bias} ->
add_bias_row(Y, Bias)
end end).
-file("src/viva_tensor/nn/attention.gleam", 443).
?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,
_ = Num_heads,
{ok, {tensor, Combined, [Seq, Num_heads * Head_dim]}}
end
).
-file("src/viva_tensor/nn/attention.gleam", 415).
?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/attention.gleam", 372).
?DOC(false).
-spec check_mha_input(binary(), viva_tensor@tensor:tensor(), integer()) -> {ok,
integer()} |
{error, viva_tensor@core@error:tensor_error()}.
check_mha_input(_, T, Embed_dim) ->
case erlang:element(3, T) of
[Seq, E] when E =:= Embed_dim ->
{ok, Seq};
_ ->
{error, {shape_mismatch, [-1, Embed_dim], erlang:element(3, T)}}
end.
-file("src/viva_tensor/nn/attention.gleam", 302).
?DOC(false).
-spec multi_head_attention_forward(
multi_head_attention(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
boolean()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
multi_head_attention_forward(Mha, Q, K, V, Is_causal) ->
gleam@result:'try'(
check_mha_input(<<"q"/utf8>>, Q, erlang:element(3, Mha)),
fun(Seq) ->
gleam@result:'try'(
check_mha_input(<<"k"/utf8>>, K, erlang:element(3, Mha)),
fun(Seq_k) ->
gleam@result:'try'(
check_mha_input(<<"v"/utf8>>, V, erlang:element(3, Mha)),
fun(Seq_v) ->
case (Seq =:= Seq_k) andalso (Seq_k =:= Seq_v) of
false ->
{error,
{shape_mismatch,
[Seq, erlang:element(3, Mha)],
erlang:element(3, K)}};
true ->
gleam@result:'try'(
linear(
Q,
erlang:element(5, Mha),
erlang:element(9, Mha)
),
fun(Q_proj) ->
gleam@result:'try'(
linear(
K,
erlang:element(6, Mha),
erlang:element(10, Mha)
),
fun(K_proj) ->
gleam@result:'try'(
linear(
V,
erlang:element(
7,
Mha
),
erlang:element(
11,
Mha
)
),
fun(V_proj) ->
gleam@result:'try'(
split_heads(
Q_proj,
Seq,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(Q_heads) ->
gleam@result:'try'(
split_heads(
K_proj,
Seq,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(
K_heads
) ->
gleam@result:'try'(
split_heads(
V_proj,
Seq,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(
V_heads
) ->
Heads_zipped = 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(
Heads_zipped,
fun(
T
) ->
{Qh@1,
Kh@1,
Vh@1} = T,
scaled_dot_product_attention(
Qh@1,
Kh@1,
Vh@1,
none,
Is_causal
)
end
),
fun(
Head_outputs
) ->
gleam@result:'try'(
concat_heads(
Head_outputs,
Seq,
erlang:element(
2,
Mha
),
erlang:element(
4,
Mha
)
),
fun(
Concat
) ->
linear(
Concat,
erlang:element(
8,
Mha
),
erlang:element(
12,
Mha
)
)
end
)
end
)
end
)
end
)
end
)
end
)
end
)
end
)
end
end
)
end
)
end
).