Current section

Files

Jump to
viva_tensor src viva_tensor@io@hf_loader.erl
Raw

src/viva_tensor@io@hf_loader.erl

-module(viva_tensor@io@hf_loader).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/io/hf_loader.gleam").
-export([load_safetensors_dict/1, load_embedding/4, load_layer_norm/3, load_multi_head_attention/4, load_feed_forward/5, load_encoder_block/6, load_transformer/7, from_safetensors_file/2]).
-export_type([hf_load_error/0, transformer_config/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 hf_load_error() :: {weight_not_found, binary()} |
{shape_mismatch, binary(), list(integer()), list(integer())} |
{io_error, binary()}.
-type transformer_config() :: {transformer_config,
integer(),
integer(),
integer(),
integer(),
integer(),
viva_tensor@nn@transformer:activation(),
boolean(),
integer()}.
-file("src/viva_tensor/io/hf_loader.gleam", 653).
?DOC(false).
-spec shape_to_string(list(integer())) -> binary().
shape_to_string(Shape) ->
<<<<"["/utf8,
(gleam@string:join(
gleam@list:map(Shape, fun erlang:integer_to_binary/1),
<<", "/utf8>>
))/binary>>/binary,
"]"/utf8>>.
-file("src/viva_tensor/io/hf_loader.gleam", 634).
?DOC(false).
-spec tensor_error_to_string(viva_tensor@core@error:tensor_error()) -> binary().
tensor_error_to_string(Err) ->
case Err of
{invalid_shape, Reason} ->
Reason;
{dtype_error, Reason@1} ->
Reason@1;
{shape_mismatch, Expected, Got} ->
<<<<<<"shape mismatch: expected "/utf8,
(shape_to_string(Expected))/binary>>/binary,
", got "/utf8>>/binary,
(shape_to_string(Got))/binary>>;
{dimension_error, Reason@2} ->
Reason@2;
{index_out_of_bounds, Idx, Size} ->
<<<<<<"index "/utf8, (erlang:integer_to_binary(Idx))/binary>>/binary,
" out of bounds for size "/utf8>>/binary,
(erlang:integer_to_binary(Size))/binary>>;
Other ->
gleam@string:inspect(Other)
end.
-file("src/viva_tensor/io/hf_loader.gleam", 136).
?DOC(false).
-spec load_safetensors_dict(binary()) -> {ok,
gleam@dict:dict(binary(), viva_tensor@tensor:tensor())} |
{error, hf_load_error()}.
load_safetensors_dict(Path) ->
case viva_tensor@io@safetensors:read(Path) of
{ok, D} ->
{ok, D};
{error, Err} ->
{error, {io_error, tensor_error_to_string(Err)}}
end.
-file("src/viva_tensor/io/hf_loader.gleam", 612).
?DOC(false).
-spec check_shape(binary(), viva_tensor@tensor:tensor(), list(integer())) -> {ok,
nil} |
{error, hf_load_error()}.
check_shape(Name, T, Expected) ->
Got = viva_tensor@tensor:shape(T),
case Got =:= Expected of
true ->
{ok, nil};
false ->
{error, {shape_mismatch, Name, Expected, Got}}
end.
-file("src/viva_tensor/io/hf_loader.gleam", 602).
?DOC(false).
-spec get_weight(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary()
) -> {ok, viva_tensor@tensor:tensor()} | {error, hf_load_error()}.
get_weight(Weights, Name) ->
case gleam_stdlib:map_get(Weights, Name) of
{ok, T} ->
{ok, T};
{error, _} ->
{error, {weight_not_found, Name}}
end.
-file("src/viva_tensor/io/hf_loader.gleam", 161).
?DOC(false).
-spec load_embedding(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary(),
integer(),
integer()
) -> {ok, viva_tensor@nn@embedding:embedding()} | {error, hf_load_error()}.
load_embedding(Weights, Prefix, Vocab_size, Embedding_dim) ->
Name = <<Prefix/binary, ".weight"/utf8>>,
gleam@result:'try'(
get_weight(Weights, Name),
fun(W) ->
gleam@result:'try'(
check_shape(Name, W, [Vocab_size, Embedding_dim]),
fun(_) -> {ok, {embedding, Vocab_size, Embedding_dim, W}} end
)
end
).
-file("src/viva_tensor/io/hf_loader.gleam", 198).
?DOC(false).
-spec load_layer_norm(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary(),
integer()
) -> {ok, viva_tensor@nn@norm:layer_norm()} | {error, hf_load_error()}.
load_layer_norm(Weights, Prefix, Num_features) ->
Scale_name = <<Prefix/binary, ".weight"/utf8>>,
Bias_name = <<Prefix/binary, ".bias"/utf8>>,
gleam@result:'try'(
get_weight(Weights, Scale_name),
fun(Scale) ->
gleam@result:'try'(
check_shape(Scale_name, Scale, [Num_features]),
fun(_) ->
gleam@result:'try'(
get_weight(Weights, Bias_name),
fun(Bias) ->
gleam@result:'try'(
check_shape(Bias_name, Bias, [Num_features]),
fun(_) ->
{ok, {layer_norm, Scale, Bias, 1.0e-5}}
end
)
end
)
end
)
end
).
-file("src/viva_tensor/io/hf_loader.gleam", 624).
?DOC(false).
-spec get_and_check(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary(),
list(integer())
) -> {ok, viva_tensor@tensor:tensor()} | {error, hf_load_error()}.
get_and_check(Weights, Name, Expected) ->
gleam@result:'try'(
get_weight(Weights, Name),
fun(T) ->
gleam@result:'try'(
check_shape(Name, T, Expected),
fun(_) -> {ok, T} end
)
end
).
-file("src/viva_tensor/io/hf_loader.gleam", 238).
?DOC(false).
-spec load_multi_head_attention(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary(),
integer(),
integer()
) -> {ok, viva_tensor@nn@attention:multi_head_attention()} |
{error, hf_load_error()}.
load_multi_head_attention(Weights, Prefix, Num_heads, Embed_dim) ->
case (Num_heads =< 0) orelse (Embed_dim =< 0) of
true ->
{error,
{io_error,
<<"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) =:= 0 of
false ->
{error,
{io_error,
<<<<<<<<"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>>}};
true ->
Head_dim = case Num_heads of
0 -> 0;
Gleam@denominator@1 -> Embed_dim div Gleam@denominator@1
end,
Weight_shape = [Embed_dim, Embed_dim],
Bias_shape = [Embed_dim],
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary, ".q_proj.weight"/utf8>>,
Weight_shape
),
fun(W_q) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary, ".q_proj.bias"/utf8>>,
Bias_shape
),
fun(B_q) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary,
".k_proj.weight"/utf8>>,
Weight_shape
),
fun(W_k) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary,
".k_proj.bias"/utf8>>,
Bias_shape
),
fun(B_k) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary,
".v_proj.weight"/utf8>>,
Weight_shape
),
fun(W_v) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary,
".v_proj.bias"/utf8>>,
Bias_shape
),
fun(B_v) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary,
".out_proj.weight"/utf8>>,
Weight_shape
),
fun(W_o) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary,
".out_proj.bias"/utf8>>,
Bias_shape
),
fun(
B_o
) ->
{ok,
{multi_head_attention,
Num_heads,
Embed_dim,
Head_dim,
W_q,
W_k,
W_v,
W_o,
{some,
B_q},
{some,
B_k},
{some,
B_v},
{some,
B_o}}}
end
)
end
)
end
)
end
)
end
)
end
)
end
)
end
)
end
end.
-file("src/viva_tensor/io/hf_loader.gleam", 345).
?DOC(false).
-spec load_feed_forward(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary(),
integer(),
integer(),
viva_tensor@nn@transformer:activation()
) -> {ok, viva_tensor@nn@transformer:feed_forward()} | {error, hf_load_error()}.
load_feed_forward(Weights, Prefix, Embed_dim, Hidden_dim, Activation) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary, ".linear1.weight"/utf8>>,
[Embed_dim, Hidden_dim]
),
fun(W1) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary, ".linear1.bias"/utf8>>,
[Hidden_dim]
),
fun(B1) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary, ".linear2.weight"/utf8>>,
[Hidden_dim, Embed_dim]
),
fun(W2) ->
gleam@result:'try'(
get_and_check(
Weights,
<<Prefix/binary, ".linear2.bias"/utf8>>,
[Embed_dim]
),
fun(B2) ->
{ok,
{feed_forward,
W1,
B1,
W2,
B2,
Activation}}
end
)
end
)
end
)
end
).
-file("src/viva_tensor/io/hf_loader.gleam", 385).
?DOC(false).
-spec load_encoder_block(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary(),
integer(),
integer(),
integer(),
viva_tensor@nn@transformer:activation()
) -> {ok, viva_tensor@nn@transformer:encoder_block()} | {error, hf_load_error()}.
load_encoder_block(
Weights,
Prefix,
Num_heads,
Embed_dim,
Hidden_dim,
Activation
) ->
gleam@result:'try'(
load_multi_head_attention(
Weights,
<<Prefix/binary, ".self_attn"/utf8>>,
Num_heads,
Embed_dim
),
fun(Mha) ->
gleam@result:'try'(
load_layer_norm(
Weights,
<<Prefix/binary, ".norm1"/utf8>>,
Embed_dim
),
fun(Norm1) ->
gleam@result:'try'(
load_layer_norm(
Weights,
<<Prefix/binary, ".norm2"/utf8>>,
Embed_dim
),
fun(Norm2) ->
gleam@result:'try'(
load_feed_forward(
Weights,
<<Prefix/binary, ".ffn"/utf8>>,
Embed_dim,
Hidden_dim,
Activation
),
fun(Ffn) ->
{ok,
{encoder_block, Mha, Ffn, Norm1, Norm2}}
end
)
end
)
end
)
end
).
-file("src/viva_tensor/io/hf_loader.gleam", 436).
?DOC(false).
-spec load_decoder_block(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
binary(),
integer(),
integer(),
integer(),
viva_tensor@nn@transformer:activation()
) -> {ok, viva_tensor@nn@transformer:decoder_block()} | {error, hf_load_error()}.
load_decoder_block(
Weights,
Prefix,
Num_heads,
Embed_dim,
Hidden_dim,
Activation
) ->
gleam@result:'try'(
load_multi_head_attention(
Weights,
<<Prefix/binary, ".self_attn"/utf8>>,
Num_heads,
Embed_dim
),
fun(Self_mha) ->
gleam@result:'try'(
load_multi_head_attention(
Weights,
<<Prefix/binary, ".cross_attn"/utf8>>,
Num_heads,
Embed_dim
),
fun(Cross_mha) ->
gleam@result:'try'(
load_layer_norm(
Weights,
<<Prefix/binary, ".norm1"/utf8>>,
Embed_dim
),
fun(Norm1) ->
gleam@result:'try'(
load_layer_norm(
Weights,
<<Prefix/binary, ".norm2"/utf8>>,
Embed_dim
),
fun(Norm2) ->
gleam@result:'try'(
load_layer_norm(
Weights,
<<Prefix/binary, ".norm3"/utf8>>,
Embed_dim
),
fun(Norm3) ->
gleam@result:'try'(
load_feed_forward(
Weights,
<<Prefix/binary,
".ffn"/utf8>>,
Embed_dim,
Hidden_dim,
Activation
),
fun(Ffn) ->
{ok,
{decoder_block,
Self_mha,
Cross_mha,
Ffn,
Norm1,
Norm2,
Norm3}}
end
)
end
)
end
)
end
)
end
)
end
).
-file("src/viva_tensor/io/hf_loader.gleam", 661).
?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/io/hf_loader.gleam", 657).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/io/hf_loader.gleam", 503).
?DOC(false).
-spec load_transformer(
gleam@dict:dict(binary(), viva_tensor@tensor:tensor()),
integer(),
integer(),
integer(),
integer(),
integer(),
viva_tensor@nn@transformer:activation()
) -> {ok, viva_tensor@nn@transformer:transformer()} | {error, hf_load_error()}.
load_transformer(
Weights,
Num_enc_layers,
Num_dec_layers,
Embed_dim,
Num_heads,
Hidden_dim,
Activation
) ->
case (Num_enc_layers < 0) orelse (Num_dec_layers < 0) of
true ->
{error,
{io_error,
<<<<<<<<"load_transformer: layer counts must be non-negative (got num_enc_layers="/utf8,
(erlang:integer_to_binary(Num_enc_layers))/binary>>/binary,
", num_dec_layers="/utf8>>/binary,
(erlang:integer_to_binary(Num_dec_layers))/binary>>/binary,
")"/utf8>>}};
false ->
Enc_indices = case Num_enc_layers =< 0 of
true ->
[];
false ->
range_int(0, Num_enc_layers - 1)
end,
Dec_indices = case Num_dec_layers =< 0 of
true ->
[];
false ->
range_int(0, Num_dec_layers - 1)
end,
gleam@result:'try'(
gleam@list:try_map(
Enc_indices,
fun(I) ->
load_encoder_block(
Weights,
<<"encoder.layers."/utf8,
(erlang:integer_to_binary(I))/binary>>,
Num_heads,
Embed_dim,
Hidden_dim,
Activation
)
end
),
fun(Encoders) ->
gleam@result:'try'(
gleam@list:try_map(
Dec_indices,
fun(I@1) ->
load_decoder_block(
Weights,
<<"decoder.layers."/utf8,
(erlang:integer_to_binary(I@1))/binary>>,
Num_heads,
Embed_dim,
Hidden_dim,
Activation
)
end
),
fun(Decoders) ->
{ok,
{transformer,
Encoders,
Decoders,
Num_enc_layers,
Num_dec_layers}}
end
)
end
)
end.
-file("src/viva_tensor/io/hf_loader.gleam", 582).
?DOC(false).
-spec from_safetensors_file(binary(), transformer_config()) -> {ok,
viva_tensor@nn@transformer:transformer()} |
{error, hf_load_error()}.
from_safetensors_file(Path, Config) ->
gleam@result:'try'(
load_safetensors_dict(Path),
fun(Weights) ->
load_transformer(
Weights,
erlang:element(2, Config),
erlang:element(3, Config),
erlang:element(4, Config),
erlang:element(5, Config),
erlang:element(6, Config),
erlang:element(7, Config)
)
end
).