-module(viva_tensor@native@inference). -compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]). -define(FILEPATH, "src/viva_tensor/native/inference.gleam"). -export([prepack_fp8_weight/1, prepack_int8_sparse_24_weight/1, prepack_int4_sparse_24_weight/1, linear_fp8/3, linear_int4_sparse/3, linear_int8_sparse/3, linear_gelu_fp8/3, linear_swiglu_fp8/4, fp8_features/1, int8_features/1, int4_features/1]). -export_type([packed_weight_fp8/0, packed_weight_int8_sparse/0, packed_weight_int4_sparse/0, bias_arg/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). -opaque packed_weight_fp8() :: {packed_weight_fp8, gleam@dynamic:dynamic_(), integer(), integer(), float()}. -opaque packed_weight_int8_sparse() :: {packed_weight_int8_sparse, gleam@dynamic:dynamic_(), integer(), integer(), list(float())}. -opaque packed_weight_int4_sparse() :: {packed_weight_int4_sparse, gleam@dynamic:dynamic_(), integer(), integer(), list(float())}. -type bias_arg() :: {bias_list, list(float())} | bias_nil. -file("src/viva_tensor/native/inference.gleam", 450). ?DOC(false). -spec join_int_list(list(integer())) -> binary(). join_int_list(Xs) -> case Xs of [] -> <<""/utf8>>; [X] -> erlang:integer_to_binary(X); [X@1 | Rest] -> <<<<(erlang:integer_to_binary(X@1))/binary, ", "/utf8>>/binary, (join_int_list(Rest))/binary>> end. -file("src/viva_tensor/native/inference.gleam", 446). ?DOC(false). -spec shape_string(list(integer())) -> binary(). shape_string(Shape) -> <<<<"["/utf8, (join_int_list(Shape))/binary>>/binary, "]"/utf8>>. -file("src/viva_tensor/native/inference.gleam", 100). ?DOC(false). -spec prepack_fp8_weight(viva_tensor@tensor:tensor()) -> {ok, packed_weight_fp8()} | {error, viva_tensor@core@error:tensor_error()}. prepack_fp8_weight(Weight) -> case viva_tensor@tensor:shape(Weight) of [In_f, Out_f] -> case viva_tensor_zig:nt_prepack_fp8( viva_tensor_inference_ffi:floats_to_fp32_binary( viva_tensor@tensor:to_list(Weight) ), [In_f, Out_f] ) of {ok, {Handle, _, _, Scale}} -> {ok, {packed_weight_fp8, Handle, In_f, Out_f, Scale}}; {error, Reason} -> {error, {dimension_error, <<"prepack_fp8_weight failed: "/utf8, Reason/binary>>}} end; Other -> {error, {dimension_error, <<"prepack_fp8_weight: expected 2-D weight, got "/utf8, (shape_string(Other))/binary>>}} end. -file("src/viva_tensor/native/inference.gleam", 131). ?DOC(false). -spec prepack_int8_sparse_24_weight(viva_tensor@tensor:tensor()) -> {ok, packed_weight_int8_sparse()} | {error, viva_tensor@core@error:tensor_error()}. prepack_int8_sparse_24_weight(Weight) -> case viva_tensor@tensor:shape(Weight) of [In_f, Out_f] -> case viva_tensor_zig:nt_prepack_int8_sparse( viva_tensor_inference_ffi:floats_to_fp32_binary( viva_tensor@tensor:to_list(Weight) ), [In_f, Out_f] ) of {ok, Handle} -> {ok, {packed_weight_int8_sparse, Handle, In_f, Out_f, []}}; {error, Reason} -> {error, {dimension_error, <<"prepack_int8_sparse_24_weight failed: "/utf8, Reason/binary>>}} end; Other -> {error, {dimension_error, <<"prepack_int8_sparse_24_weight: expected 2-D weight, got "/utf8, (shape_string(Other))/binary>>}} end. -file("src/viva_tensor/native/inference.gleam", 168). ?DOC(false). -spec prepack_int4_sparse_24_weight(viva_tensor@tensor:tensor()) -> {ok, packed_weight_int4_sparse()} | {error, viva_tensor@core@error:tensor_error()}. prepack_int4_sparse_24_weight(Weight) -> case viva_tensor@tensor:shape(Weight) of [In_f, Out_f] -> case viva_tensor_zig:nt_prepack_int4_sparse( viva_tensor_inference_ffi:floats_to_fp32_binary( viva_tensor@tensor:to_list(Weight) ), [In_f, Out_f] ) of {ok, Handle} -> {ok, {packed_weight_int4_sparse, Handle, In_f, Out_f, []}}; {error, Reason} -> {error, {dimension_error, <<"prepack_int4_sparse_24_weight failed: "/utf8, Reason/binary>>}} end; Other -> {error, {dimension_error, <<"prepack_int4_sparse_24_weight: expected 2-D weight, got "/utf8, (shape_string(Other))/binary>>}} end. -file("src/viva_tensor/native/inference.gleam", 433). ?DOC(false). -spec make_2d_tensor(list(float()), integer(), integer()) -> viva_tensor@tensor:tensor(). make_2d_tensor(Data, Rows, Cols) -> case viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Data), [Rows, Cols] ) of {ok, T} -> T; {error, _} -> viva_tensor@tensor:from_list(Data) end. -file("src/viva_tensor/native/inference.gleam", 426). ?DOC(false). -spec optional_tensor_to_bias_arg( gleam@option:option(viva_tensor@tensor:tensor()) ) -> bias_arg(). optional_tensor_to_bias_arg(Bias) -> case Bias of {some, B} -> {bias_list, viva_tensor@tensor:to_list(B)}; none -> bias_nil end. -file("src/viva_tensor/native/inference.gleam", 209). ?DOC(false). -spec linear_fp8( viva_tensor@tensor:tensor(), packed_weight_fp8(), gleam@option:option(viva_tensor@tensor:tensor()) ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. linear_fp8(Input, Weight, Bias) -> case viva_tensor@tensor:shape(Input) of [Batch, In_f] when In_f =:= erlang:element(3, Weight) -> Bias_data = optional_tensor_to_bias_arg(Bias), _ = Batch, _ = In_f, Input_bin = viva_tensor_inference_ffi:floats_to_fp16_binary( viva_tensor@tensor:to_list(Input) ), case viva_tensor_zig:nt_linear_fp8( Input_bin, erlang:element(2, Weight), Bias_data, 1 ) of {ok, Out_bin} -> Out_data = viva_tensor_inference_ffi:fp16_binary_to_floats( Out_bin ), {ok, make_2d_tensor( Out_data, Batch, erlang:element(4, Weight) )}; {error, Reason} -> {error, {dimension_error, <<"linear_fp8 failed: "/utf8, Reason/binary>>}} end; Other -> {error, {dimension_error, <<<<<<<<"linear_fp8: input feature dim mismatch (got "/utf8, (shape_string(Other))/binary>>/binary, ", weight expects "/utf8>>/binary, (erlang:integer_to_binary(erlang:element(3, Weight)))/binary>>/binary, ")"/utf8>>}} end. -file("src/viva_tensor/native/inference.gleam", 244). ?DOC(false). -spec linear_int4_sparse( viva_tensor@tensor:tensor(), packed_weight_int4_sparse(), gleam@option:option(viva_tensor@tensor:tensor()) ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. linear_int4_sparse(Input, Weight, Bias) -> case viva_tensor@tensor:shape(Input) of [Batch, In_f] when In_f =:= erlang:element(3, Weight) -> Bias_data = optional_tensor_to_bias_arg(Bias), _ = Batch, _ = In_f, Input_bin = viva_tensor_inference_ffi:floats_to_fp16_binary( viva_tensor@tensor:to_list(Input) ), case viva_tensor_zig:nt_linear_int4_sparse( Input_bin, erlang:element(2, Weight), Bias_data ) of {ok, Out_bin} -> Out_data = viva_tensor_inference_ffi:fp16_binary_to_floats( Out_bin ), {ok, make_2d_tensor( Out_data, Batch, erlang:element(4, Weight) )}; {error, Reason} -> {error, {dimension_error, <<"linear_int4_sparse failed: "/utf8, Reason/binary>>}} end; Other -> {error, {dimension_error, <<<<<<<<"linear_int4_sparse: input feature dim mismatch (got "/utf8, (shape_string(Other))/binary>>/binary, ", weight expects "/utf8>>/binary, (erlang:integer_to_binary(erlang:element(3, Weight)))/binary>>/binary, ")"/utf8>>}} end. -file("src/viva_tensor/native/inference.gleam", 278). ?DOC(false). -spec linear_int8_sparse( viva_tensor@tensor:tensor(), packed_weight_int8_sparse(), gleam@option:option(viva_tensor@tensor:tensor()) ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. linear_int8_sparse(Input, Weight, Bias) -> case viva_tensor@tensor:shape(Input) of [Batch, In_f] when In_f =:= erlang:element(3, Weight) -> Bias_data = optional_tensor_to_bias_arg(Bias), _ = Batch, _ = In_f, Input_bin = viva_tensor_inference_ffi:floats_to_fp16_binary( viva_tensor@tensor:to_list(Input) ), case viva_tensor_zig:nt_linear_int8_sparse( Input_bin, erlang:element(2, Weight), Bias_data ) of {ok, Out_bin} -> Out_data = viva_tensor_inference_ffi:fp16_binary_to_floats( Out_bin ), {ok, make_2d_tensor( Out_data, Batch, erlang:element(4, Weight) )}; {error, Reason} -> {error, {dimension_error, <<"linear_int8_sparse failed: "/utf8, Reason/binary>>}} end; Other -> {error, {dimension_error, <<<<<<<<"linear_int8_sparse: input feature dim mismatch (got "/utf8, (shape_string(Other))/binary>>/binary, ", weight expects "/utf8>>/binary, (erlang:integer_to_binary(erlang:element(3, Weight)))/binary>>/binary, ")"/utf8>>}} end. -file("src/viva_tensor/native/inference.gleam", 313). ?DOC(false). -spec linear_gelu_fp8( viva_tensor@tensor:tensor(), packed_weight_fp8(), gleam@option:option(viva_tensor@tensor:tensor()) ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. linear_gelu_fp8(Input, Weight, Bias) -> case viva_tensor@tensor:shape(Input) of [Batch, In_f] when In_f =:= erlang:element(3, Weight) -> Bias_data = optional_tensor_to_bias_arg(Bias), _ = Batch, _ = In_f, Input_bin = viva_tensor_inference_ffi:floats_to_fp16_binary( viva_tensor@tensor:to_list(Input) ), case viva_tensor_zig:nt_linear_gelu_fp8( Input_bin, erlang:element(2, Weight), Bias_data, 36 ) of {ok, Out_bin} -> Out_data = viva_tensor_inference_ffi:fp16_binary_to_floats( Out_bin ), {ok, make_2d_tensor( Out_data, Batch, erlang:element(4, Weight) )}; {error, Reason} -> {error, {dimension_error, <<"linear_gelu_fp8 failed: "/utf8, Reason/binary>>}} end; Other -> {error, {dimension_error, <<<<<<<<"linear_gelu_fp8: input feature dim mismatch (got "/utf8, (shape_string(Other))/binary>>/binary, ", weight expects "/utf8>>/binary, (erlang:integer_to_binary(erlang:element(3, Weight)))/binary>>/binary, ")"/utf8>>}} end. -file("src/viva_tensor/native/inference.gleam", 352). ?DOC(false). -spec linear_swiglu_fp8( viva_tensor@tensor:tensor(), packed_weight_fp8(), packed_weight_fp8(), gleam@option:option(viva_tensor@tensor:tensor()) ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. linear_swiglu_fp8(Input, Gate_weight, Up_weight, Bias) -> case {viva_tensor@tensor:shape(Input), erlang:element(3, Gate_weight) =:= erlang:element(3, Up_weight), erlang:element(4, Gate_weight) =:= erlang:element(4, Up_weight)} of {[Batch, In_f], true, true} when In_f =:= erlang:element(3, Gate_weight) -> Bias_data = optional_tensor_to_bias_arg(Bias), Input_data = viva_tensor@tensor:to_list(Input), case viva_tensor_zig:nt_linear_swiglu_fp8( Input_data, [Batch, In_f], erlang:element(2, Gate_weight), erlang:element(2, Up_weight), Bias_data ) of {ok, Out_data} -> {ok, make_2d_tensor( Out_data, Batch, erlang:element(4, Gate_weight) )}; {error, Reason} -> {error, {dimension_error, <<"linear_swiglu_fp8 failed: "/utf8, Reason/binary>>}} end; {_, false, _} -> {error, {dimension_error, <<"linear_swiglu_fp8: gate/up in_features mismatch"/utf8>>}}; {_, _, false} -> {error, {dimension_error, <<"linear_swiglu_fp8: gate/up out_features mismatch"/utf8>>}}; {Other_shape, _, _} -> {error, {dimension_error, <<"linear_swiglu_fp8: bad input shape "/utf8, (shape_string(Other_shape))/binary>>}} end. -file("src/viva_tensor/native/inference.gleam", 398). ?DOC(false). -spec fp8_features(packed_weight_fp8()) -> {integer(), integer()}. fp8_features(W) -> {erlang:element(3, W), erlang:element(4, W)}. -file("src/viva_tensor/native/inference.gleam", 403). ?DOC(false). -spec int8_features(packed_weight_int8_sparse()) -> {integer(), integer()}. int8_features(W) -> {erlang:element(3, W), erlang:element(4, W)}. -file("src/viva_tensor/native/inference.gleam", 408). ?DOC(false). -spec int4_features(packed_weight_int4_sparse()) -> {integer(), integer()}. int4_features(W) -> {erlang:element(3, W), erlang:element(4, W)}.