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@native@inference.erl
-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)}.