-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 = <>, 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 = <>, Bias_name = <>, 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, <>, Weight_shape ), fun(W_q) -> gleam@result:'try'( get_and_check( Weights, <>, Bias_shape ), fun(B_q) -> gleam@result:'try'( get_and_check( Weights, <>, Weight_shape ), fun(W_k) -> gleam@result:'try'( get_and_check( Weights, <>, Bias_shape ), fun(B_k) -> gleam@result:'try'( get_and_check( Weights, <>, Weight_shape ), fun(W_v) -> gleam@result:'try'( get_and_check( Weights, <>, Bias_shape ), fun(B_v) -> gleam@result:'try'( get_and_check( Weights, <>, Weight_shape ), fun(W_o) -> gleam@result:'try'( get_and_check( Weights, <>, 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, <>, [Embed_dim, Hidden_dim] ), fun(W1) -> gleam@result:'try'( get_and_check( Weights, <>, [Hidden_dim] ), fun(B1) -> gleam@result:'try'( get_and_check( Weights, <>, [Hidden_dim, Embed_dim] ), fun(W2) -> gleam@result:'try'( get_and_check( Weights, <>, [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, <>, Num_heads, Embed_dim ), fun(Mha) -> gleam@result:'try'( load_layer_norm( Weights, <>, Embed_dim ), fun(Norm1) -> gleam@result:'try'( load_layer_norm( Weights, <>, Embed_dim ), fun(Norm2) -> gleam@result:'try'( load_feed_forward( Weights, <>, 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, <>, Num_heads, Embed_dim ), fun(Self_mha) -> gleam@result:'try'( load_multi_head_attention( Weights, <>, Num_heads, Embed_dim ), fun(Cross_mha) -> gleam@result:'try'( load_layer_norm( Weights, <>, Embed_dim ), fun(Norm1) -> gleam@result:'try'( load_layer_norm( Weights, <>, Embed_dim ), fun(Norm2) -> gleam@result:'try'( load_layer_norm( Weights, <>, Embed_dim ), fun(Norm3) -> gleam@result:'try'( load_feed_forward( Weights, <>, 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", 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 -> gleam@list:range(0, Num_enc_layers - 1) end, Dec_indices = case Num_dec_layers =< 0 of true -> []; false -> gleam@list:range(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 ).