Packages
Tensor library for Gleam/BEAM with a pure Gleam API, zero-copy views, and optional native acceleration
Current section
Files
Jump to
Current section
Files
src/viva_tensor@nn@moe.erl
-module(viva_tensor@nn@moe).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/nn/moe.gleam").
-export([router_init/3, expert_distribution/3, compute_load_balance_loss/3, router_route/2, moe_block_init/4, moe_block_forward/2]).
-export_type([router/0, moe_block/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 router() :: {router,
viva_tensor@tensor:tensor(),
integer(),
integer(),
float()}.
-type moe_block() :: {moe_block,
router(),
list(viva_tensor@tensor:tensor()),
list(viva_tensor@tensor:tensor())}.
-file("src/viva_tensor/nn/moe.gleam", 70).
?DOC(false).
-spec router_init(integer(), integer(), integer()) -> router().
router_init(Embed_dim, Num_experts, Top_k) ->
{router,
viva_tensor@tensor:zeros([Embed_dim, Num_experts]),
Num_experts,
Top_k,
+0.0}.
-file("src/viva_tensor/nn/moe.gleam", 515).
?DOC(false).
-spec increment_at(list(integer()), integer()) -> list(integer()).
increment_at(Xs, Idx) ->
gleam@list:index_map(Xs, fun(V, I) -> case I =:= Idx of
true ->
V + 1;
false ->
V
end end).
-file("src/viva_tensor/nn/moe.gleam", 528).
?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/moe.gleam", 524).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/nn/moe.gleam", 459).
?DOC(false).
-spec expert_distribution(
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
integer()
) -> {ok, {viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
expert_distribution(Router_probs, Expert_assignments, Num_experts) ->
case {erlang:element(3, Router_probs),
erlang:element(3, Expert_assignments)} of
{[Tokens_p, Ne], [Tokens_a, _]} when Ne =:= Num_experts ->
case Tokens_p =:= Tokens_a of
false ->
{error,
{shape_mismatch,
erlang:element(3, Router_probs),
erlang:element(3, Expert_assignments)}};
true ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Router_probs),
fun(Probs_data) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(
Expert_assignments
),
fun(Ids_data) ->
Prob_rows = gleam@list:sized_chunk(
Probs_data,
Num_experts
),
Importance_data = begin
_pipe = range_int(0, Num_experts - 1),
gleam@list:map(
_pipe,
fun(I) ->
gleam@list:fold(
Prob_rows,
+0.0,
fun(Acc, Row) ->
Acc + begin
_pipe@1 = Row,
_pipe@2 = gleam@list:drop(
_pipe@1,
I
),
_pipe@3 = gleam@list:first(
_pipe@2
),
gleam@result:unwrap(
_pipe@3,
+0.0
)
end
end
)
end
)
end,
Counts = gleam@list:fold(
Ids_data,
gleam@list:repeat(0, Num_experts),
fun(Acc@1, Id_f) ->
Id = erlang:round(Id_f),
case (Id >= 0) andalso (Id < Num_experts) of
false ->
Acc@1;
true ->
increment_at(Acc@1, Id)
end
end
),
Load_data = gleam@list:map(
Counts,
fun erlang:float/1
),
{ok,
{{tensor,
Importance_data,
[Num_experts]},
{tensor, Load_data, [Num_experts]}}}
end
)
end
)
end;
{_, _} ->
{error,
{shape_mismatch,
[-1, Num_experts],
erlang:element(3, Router_probs)}}
end.
-file("src/viva_tensor/nn/moe.gleam", 405).
?DOC(false).
-spec compute_load_balance_loss(
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
integer()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
compute_load_balance_loss(Router_probs, Expert_assignments, Num_experts) ->
gleam@result:'try'(
expert_distribution(Router_probs, Expert_assignments, Num_experts),
fun(_use0) ->
{Importance, Load} = _use0,
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Importance),
fun(Importance_data) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Load),
fun(Load_data) ->
Num_tokens = case erlang:element(3, Router_probs) of
[N, _] ->
N;
_ ->
0
end,
Top_k = case erlang:element(3, Expert_assignments) of
[_, K] ->
K;
_ ->
0
end,
Total_assignments = Num_tokens * Top_k,
case (Num_tokens =< 0) orelse (Total_assignments =< 0) of
true ->
{ok, {tensor, [+0.0], []}};
false ->
N_f = erlang:float(Num_tokens),
Total_f = erlang:float(Total_assignments),
Dot = begin
_pipe = gleam@list:zip(
Importance_data,
Load_data
),
gleam@list:fold(
_pipe,
+0.0,
fun(Acc, Pair) ->
{Imp, Ld} = Pair,
Acc + ((case N_f of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> Imp / Gleam@denominator
end) * (case Total_f of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> Ld / Gleam@denominator@1
end))
end
)
end,
Loss = erlang:float(Num_experts) * Dot,
{ok, {tensor, [Loss], []}}
end
end
)
end
)
end
).
-file("src/viva_tensor/nn/moe.gleam", 194).
?DOC(false).
-spec softmax_row(list(float())) -> list(float()).
softmax_row(Values) ->
case Values of
[] ->
[];
[First | Rest] ->
Max_v = gleam@list:fold(
Rest,
First,
fun(Acc, V) -> gleam@float:max(Acc, V) end
),
Shifted = gleam@list:map(Values, fun(V@1) -> V@1 - Max_v end),
Exps = gleam@list:map(Shifted, fun math:exp/1),
Sum_exp = gleam@list:fold(
Exps,
+0.0,
fun(Acc@1, V@2) -> Acc@1 + V@2 end
),
case Sum_exp > +0.0 of
true ->
gleam@list:map(Exps, fun(V@3) -> case Sum_exp of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> V@3 / Gleam@denominator
end end);
false ->
gleam@list:map(Exps, fun(_) -> +0.0 end)
end
end.
-file("src/viva_tensor/nn/moe.gleam", 182).
?DOC(false).
-spec top_k_row(list(float()), integer()) -> list({integer(), float()}).
top_k_row(Row, K) ->
Indexed = gleam@list:index_map(Row, fun(V, I) -> {I, V} end),
Sorted = gleam@list:sort(
Indexed,
fun(A, B) ->
{_, Va} = A,
{_, Vb} = B,
gleam@float:compare(Vb, Va)
end
),
gleam@list:take(Sorted, K).
-file("src/viva_tensor/nn/moe.gleam", 169).
?DOC(false).
-spec add_noise(viva_tensor@tensor:tensor(), float()) -> viva_tensor@tensor:tensor().
add_noise(T, Std) ->
Noise = viva_tensor@nn@init:normal(erlang:element(3, T), +0.0, Std),
case viva_tensor@tensor:add(T, Noise) of
{ok, Sum} ->
Sum;
{error, _} ->
T
end.
-file("src/viva_tensor/nn/moe.gleam", 98).
?DOC(false).
-spec router_route(router(), viva_tensor@tensor:tensor()) -> {ok,
{viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
router_route(Router, Tokens) ->
case {erlang:element(3, Tokens),
erlang:element(3, erlang:element(2, Router))} of
{[Num_tokens, Embed_dim], [Gate_in, Num_experts]} when (Embed_dim =:= Gate_in) andalso (Num_experts =:= erlang:element(
3,
Router
)) ->
case (erlang:element(4, Router) > 0) andalso (erlang:element(
4,
Router
)
=< erlang:element(3, Router)) of
false ->
{error,
{invalid_shape,
<<<<<<<<"router_route: top_k must be in 1..="/utf8,
(erlang:integer_to_binary(
erlang:element(3, Router)
))/binary>>/binary,
" (got "/utf8>>/binary,
(erlang:integer_to_binary(
erlang:element(4, Router)
))/binary>>/binary,
")"/utf8>>}};
true ->
gleam@result:'try'(
viva_tensor@tensor:matmul(
Tokens,
erlang:element(2, Router)
),
fun(Logits) ->
Noisy_logits = case erlang:element(5, Router) > +0.0 of
true ->
add_noise(Logits, erlang:element(5, Router));
false ->
Logits
end,
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Noisy_logits),
fun(Logits_data) ->
Rows = gleam@list:sized_chunk(
Logits_data,
Num_experts
),
Top_per_row = gleam@list:map(
Rows,
fun(Row) ->
top_k_row(
Row,
erlang:element(4, Router)
)
end
),
Ids_data = gleam@list:flat_map(
Top_per_row,
fun(Picks) ->
gleam@list:map(
Picks,
fun(P) ->
{Idx, _} = P,
erlang:float(Idx)
end
)
end
),
Expert_ids = {tensor,
Ids_data,
[Num_tokens, erlang:element(4, Router)]},
Weights_data = gleam@list:flat_map(
Top_per_row,
fun(Picks@1) ->
Top_logits = gleam@list:map(
Picks@1,
fun(P@1) ->
erlang:element(2, P@1)
end
),
softmax_row(Top_logits)
end
),
Expert_weights = {tensor,
Weights_data,
[Num_tokens, erlang:element(4, Router)]},
gleam@result:'try'(
viva_tensor@nn@activations:softmax(
Noisy_logits,
1
),
fun(Router_probs) ->
gleam@result:'try'(
compute_load_balance_loss(
Router_probs,
Expert_ids,
erlang:element(3, Router)
),
fun(Aux_loss) ->
{ok,
{Expert_ids,
Expert_weights,
Aux_loss}}
end
)
end
)
end
)
end
)
end;
{[_, _], [_, _]} ->
{error,
{shape_mismatch,
erlang:element(3, erlang:element(2, Router)),
erlang:element(3, Tokens)}};
{_, _} ->
{error,
{invalid_shape,
<<"router_route: tokens must be rank-2 [tokens, embed_dim] and gate must be [embed_dim, num_experts]"/utf8>>}}
end.
-file("src/viva_tensor/nn/moe.gleam", 232).
?DOC(false).
-spec moe_block_init(integer(), integer(), integer(), integer()) -> {ok,
moe_block()} |
{error, viva_tensor@core@error:tensor_error()}.
moe_block_init(Embed_dim, Hidden_dim, Num_experts, Top_k) ->
case ((Embed_dim =< 0) orelse (Hidden_dim =< 0)) orelse (Num_experts =< 0) of
true ->
{error,
{invalid_shape,
<<"moe_block_init: embed_dim, hidden_dim, and num_experts must be > 0"/utf8>>}};
false ->
case (Top_k =< 0) orelse (Top_k > Num_experts) of
true ->
{error,
{invalid_shape,
<<<<<<<<"moe_block_init: top_k must satisfy 1 <= top_k <= num_experts (got top_k="/utf8,
(erlang:integer_to_binary(Top_k))/binary>>/binary,
", num_experts="/utf8>>/binary,
(erlang:integer_to_binary(Num_experts))/binary>>/binary,
")"/utf8>>}};
false ->
W1s = begin
_pipe = gleam@list:repeat(nil, Num_experts),
gleam@list:map(
_pipe,
fun(_) ->
viva_tensor@tensor:zeros(
[Embed_dim, Hidden_dim]
)
end
)
end,
W2s = begin
_pipe@1 = gleam@list:repeat(nil, Num_experts),
gleam@list:map(
_pipe@1,
fun(_) ->
viva_tensor@tensor:zeros(
[Hidden_dim, Embed_dim]
)
end
)
end,
{ok,
{moe_block,
router_init(Embed_dim, Num_experts, Top_k),
W1s,
W2s}}
end
end.
-file("src/viva_tensor/nn/moe.gleam", 368).
?DOC(false).
-spec list_at(list(AFCX), integer()) -> {ok, AFCX} | {error, nil}.
list_at(Xs, Idx) ->
case Idx < 0 of
true ->
{error, nil};
false ->
_pipe = Xs,
_pipe@1 = gleam@list:drop(_pipe, Idx),
gleam@list:first(_pipe@1)
end.
-file("src/viva_tensor/nn/moe.gleam", 337).
?DOC(false).
-spec process_token_row(
list(float()),
list(float()),
list(float()),
list(viva_tensor@tensor:tensor()),
list(viva_tensor@tensor:tensor()),
integer()
) -> {ok, list(float())} | {error, viva_tensor@core@error:tensor_error()}.
process_token_row(Row, Ids, Weights, Expert_w1, Expert_w2, Embed_dim) ->
Pairs = gleam@list:zip(Ids, Weights),
gleam@list:try_fold(
Pairs,
gleam@list:repeat(+0.0, Embed_dim),
fun(Acc, Pair) ->
{Id_f, Weight} = Pair,
Id = erlang:round(Id_f),
case {list_at(Expert_w1, Id), list_at(Expert_w2, Id)} of
{{ok, W1}, {ok, W2}} ->
Row_tensor = {tensor, Row, [1, Embed_dim]},
gleam@result:'try'(
viva_tensor@tensor:matmul(Row_tensor, W1),
fun(H) ->
Activated = viva_tensor@nn@activations:relu(H),
gleam@result:'try'(
viva_tensor@tensor:matmul(Activated, W2),
fun(O) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(O),
fun(O_data) ->
{ok,
gleam@list:map2(
Acc,
O_data,
fun(A, V) ->
A + (Weight * V)
end
)}
end
)
end
)
end
);
{_, _} ->
{error,
{invalid_shape,
<<<<"moe_block_forward: expert id "/utf8,
(erlang:integer_to_binary(Id))/binary>>/binary,
" out of range"/utf8>>}}
end
end
).
-file("src/viva_tensor/nn/moe.gleam", 288).
?DOC(false).
-spec moe_block_forward(moe_block(), viva_tensor@tensor:tensor()) -> {ok,
{viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor()}} |
{error, viva_tensor@core@error:tensor_error()}.
moe_block_forward(Block, Tokens) ->
case erlang:element(3, Tokens) of
[Num_tokens, Embed_dim] ->
gleam@result:'try'(
router_route(erlang:element(2, Block), Tokens),
fun(_use0) ->
{Expert_ids, Expert_weights, Aux_loss} = _use0,
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Tokens),
fun(Token_data) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(Expert_ids),
fun(Ids_data) ->
gleam@result:'try'(
viva_tensor@tensor:try_to_list(
Expert_weights
),
fun(Weight_data) ->
Token_rows = gleam@list:sized_chunk(
Token_data,
Embed_dim
),
Id_rows = gleam@list:sized_chunk(
Ids_data,
erlang:element(
4,
erlang:element(2, Block)
)
),
Weight_rows = gleam@list:sized_chunk(
Weight_data,
erlang:element(
4,
erlang:element(2, Block)
)
),
gleam@result:'try'(
begin
_pipe = gleam@list:zip(
Token_rows,
gleam@list:zip(
Id_rows,
Weight_rows
)
),
gleam@list:try_map(
_pipe,
fun(Triple) ->
{Row, Rest} = Triple,
{Row_ids,
Row_weights} = Rest,
process_token_row(
Row,
Row_ids,
Row_weights,
erlang:element(
3,
Block
),
erlang:element(
4,
Block
),
Embed_dim
)
end
)
end,
fun(Out_rows) ->
Out_data = lists:append(
Out_rows
),
Out = {tensor,
Out_data,
[Num_tokens, Embed_dim]},
_ = Num_tokens,
{ok, {Out, Aux_loss}}
end
)
end
)
end
)
end
)
end
);
_ ->
{error,
{invalid_shape,
<<"moe_block_forward: tokens must be rank-2 [tokens, embed_dim]"/utf8>>}}
end.