Current section

Files

Jump to
viva_tensor src viva_tensor@quant@qat.erl
Raw

src/viva_tensor@quant@qat.erl

-module(viva_tensor@quant@qat).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/quant/qat.gleam").
-export([observe/2, fake_quant_forward/3, fake_quant_backward/4, compute_per_channel_scales/3, qat_linear_init/4, qat_linear_forward/2, qat_linear_calibrate/2]).
-export_type([quant_config/0, quant_stats/0, qat_linear/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 quant_config() :: {quant_config,
integer(),
boolean(),
boolean(),
integer()}.
-type quant_stats() :: {quant_stats,
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()}.
-type qat_linear() :: {qat_linear,
viva_tensor@tensor:tensor(),
gleam@option:option(viva_tensor@tensor:tensor()),
quant_stats(),
quant_config(),
gleam@option:option(quant_stats()),
quant_config()}.
-file("src/viva_tensor/quant/qat.gleam", 90).
?DOC(false).
-spec pow2(integer()) -> integer().
pow2(Exp) ->
case Exp of
E when E =< 0 ->
1;
_ ->
2 * pow2(Exp - 1)
end.
-file("src/viva_tensor/quant/qat.gleam", 77).
?DOC(false).
-spec quant_range(quant_config()) -> {float(), float()}.
quant_range(Config) ->
case erlang:element(3, Config) of
true ->
Qmax = erlang:float(pow2(erlang:element(2, Config) - 1) - 1),
{+0.0 - Qmax, Qmax};
false ->
Qmax@1 = erlang:float(pow2(erlang:element(2, Config)) - 1),
{+0.0, Qmax@1}
end.
-file("src/viva_tensor/quant/qat.gleam", 204).
?DOC(false).
-spec float_to_int(float()) -> integer().
float_to_int(F) ->
erlang:round(F).
-file("src/viva_tensor/quant/qat.gleam", 171).
?DOC(false).
-spec compute_scale_zp(float(), float(), float(), float(), boolean()) -> {float(),
integer()}.
compute_scale_zp(Min_v, Max_v, Qmin, Qmax, Symmetric) ->
case Symmetric of
true ->
Abs_max = gleam@float:max(
gleam@float:absolute_value(Min_v),
gleam@float:absolute_value(Max_v)
),
Scale = case Abs_max > +0.0 of
true ->
case Qmax of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> Abs_max / Gleam@denominator
end;
false ->
1.0
end,
{Scale, 0};
false ->
Range = Max_v - Min_v,
Q_range = Qmax - Qmin,
Scale@1 = case Range > +0.0 of
true ->
case Q_range of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> Range / Gleam@denominator@1
end;
false ->
1.0
end,
Zp_float = Qmin - (case Scale@1 of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@2 -> Min_v / Gleam@denominator@2
end),
Zp_rounded = erlang:round(Zp_float),
Zp_clamped = gleam@int:max(
gleam@int:min(Zp_rounded, float_to_int(Qmax)),
float_to_int(Qmin)
),
{Scale@1, Zp_clamped}
end.
-file("src/viva_tensor/quant/qat.gleam", 208).
?DOC(false).
-spec min_max(list(float())) -> {float(), float()}.
min_max(Values) ->
case Values of
[] ->
{+0.0, +0.0};
[First | Rest] ->
gleam@list:fold(
Rest,
{First, First},
fun(Acc, V) ->
{Lo, Hi} = Acc,
New_lo = case V < Lo of
true ->
V;
false ->
Lo
end,
New_hi = case V > Hi of
true ->
V;
false ->
Hi
end,
{New_lo, New_hi}
end
)
end.
-file("src/viva_tensor/quant/qat.gleam", 588).
?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/quant/qat.gleam", 584).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/quant/qat.gleam", 265).
?DOC(false).
-spec product_of(list(integer())) -> integer().
product_of(Dims) ->
gleam@list:fold(Dims, 1, fun(Acc, D) -> Acc * D end).
-file("src/viva_tensor/quant/qat.gleam", 230).
?DOC(false).
-spec split_along_axis(list(float()), list(integer()), integer()) -> {ok,
list(list(float()))} |
{error, viva_tensor@core@error:tensor_error()}.
split_along_axis(Data, Shape, Axis) ->
C = case gleam@list:drop(Shape, Axis) of
[Size | _] ->
Size;
[] ->
0
end,
case C =< 0 of
true ->
{error,
{invalid_shape, <<"observe: channel axis has size 0"/utf8>>}};
false ->
Outer = product_of(gleam@list:take(Shape, Axis)),
Inner = product_of(gleam@list:drop(Shape, Axis + 1)),
Buckets = gleam@list:repeat([], C),
Indexed = begin
_pipe = Data,
gleam@list:index_map(
_pipe,
fun(Value, Idx) ->
Channel = case C of
0 -> 0;
Gleam@denominator@1 -> case Inner of
0 -> 0;
Gleam@denominator -> Idx div Gleam@denominator
end rem Gleam@denominator@1
end,
_ = Outer,
{Channel, Value}
end
)
end,
Filled = begin
_pipe@1 = range_int(0, C - 1),
gleam@list:map(_pipe@1, fun(Ch) -> _pipe@2 = Indexed,
_pipe@3 = gleam@list:filter(
_pipe@2,
fun(Pair) -> erlang:element(1, Pair) =:= Ch end
),
gleam@list:map(
_pipe@3,
fun(Pair@1) -> erlang:element(2, Pair@1) end
) end)
end,
_ = Buckets,
{ok, Filled}
end.
-file("src/viva_tensor/quant/qat.gleam", 139).
?DOC(false).
-spec observe_per_channel(
viva_tensor@tensor:tensor(),
list(float()),
quant_config()
) -> {ok, quant_stats()} | {error, viva_tensor@core@error:tensor_error()}.
observe_per_channel(Input, Data, Config) ->
Shape = viva_tensor@tensor:shape(Input),
Rank = erlang:length(Shape),
Axis = erlang:element(5, Config),
case (Axis < 0) orelse (Axis >= Rank) of
true ->
{error,
{dimension_error,
<<"observe: channel_axis out of bounds for tensor rank"/utf8>>}};
false ->
gleam@result:'try'(
split_along_axis(Data, Shape, Axis),
fun(Channels) ->
{Qmin, Qmax} = quant_range(Config),
Pairs = gleam@list:map(
Channels,
fun(Values) ->
{Min_v, Max_v} = min_max(Values),
compute_scale_zp(
Min_v,
Max_v,
Qmin,
Qmax,
erlang:element(3, Config)
)
end
),
Scales = gleam@list:map(
Pairs,
fun(P) -> erlang:element(1, P) end
),
Zps = gleam@list:map(
Pairs,
fun(P@1) -> erlang:float(erlang:element(2, P@1)) end
),
C = erlang:length(Scales),
{ok,
{quant_stats, {tensor, Scales, [C]}, {tensor, Zps, [C]}}}
end
)
end.
-file("src/viva_tensor/quant/qat.gleam", 125).
?DOC(false).
-spec observe_tensor_wide(list(float()), quant_config()) -> {ok, quant_stats()} |
{error, viva_tensor@core@error:tensor_error()}.
observe_tensor_wide(Data, Config) ->
{Qmin, Qmax} = quant_range(Config),
{Min_v, Max_v} = min_max(Data),
{Scale, Zp} = compute_scale_zp(
Min_v,
Max_v,
Qmin,
Qmax,
erlang:element(3, Config)
),
{ok,
{quant_stats, {tensor, [Scale], [1]}, {tensor, [erlang:float(Zp)], [1]}}}.
-file("src/viva_tensor/quant/qat.gleam", 109).
?DOC(false).
-spec observe(viva_tensor@tensor:tensor(), quant_config()) -> {ok,
quant_stats()} |
{error, viva_tensor@core@error:tensor_error()}.
observe(Input, Config) ->
Data = viva_tensor@tensor:to_list(Input),
case Data of
[] ->
{error, {invalid_shape, <<"observe: empty tensor"/utf8>>}};
_ ->
case erlang:element(4, Config) of
false ->
observe_tensor_wide(Data, Config);
true ->
observe_per_channel(Input, Data, Config)
end
end.
-file("src/viva_tensor/quant/qat.gleam", 414).
?DOC(false).
-spec nth(list(float()), integer(), float()) -> float().
nth(Values, Index, Default) ->
case gleam@list:drop(Values, Index) of
[V | _] ->
V;
[] ->
Default
end.
-file("src/viva_tensor/quant/qat.gleam", 364).
?DOC(false).
-spec broadcast_per_element(
list(integer()),
list(float()),
list(float()),
quant_config()
) -> {ok, list({float(), float()})} |
{error, viva_tensor@core@error:tensor_error()}.
broadcast_per_element(Shape, Scales, Zps, Config) ->
Total = product_of(Shape),
case erlang:element(4, Config) of
false ->
Scale = case Scales of
[S | _] ->
S;
[] ->
1.0
end,
Zp = case Zps of
[Z | _] ->
Z;
[] ->
+0.0
end,
{ok, gleam@list:repeat({Scale, Zp}, Total)};
true ->
Rank = erlang:length(Shape),
Axis = erlang:element(5, Config),
case (Axis < 0) orelse (Axis >= Rank) of
true ->
{error,
{dimension_error,
<<"fake_quant: channel_axis out of bounds for tensor rank"/utf8>>}};
false ->
C = case gleam@list:drop(Shape, Axis) of
[Size | _] ->
Size;
[] ->
0
end,
Inner = product_of(gleam@list:drop(Shape, Axis + 1)),
Scale_arr = Scales,
Zp_arr = Zps,
Result = begin
_pipe = range_int(0, Total - 1),
gleam@list:map(
_pipe,
fun(Idx) ->
Channel = case C of
0 -> 0;
Gleam@denominator@1 -> case Inner of
0 -> 0;
Gleam@denominator -> Idx div Gleam@denominator
end rem Gleam@denominator@1
end,
S@1 = nth(Scale_arr, Channel, 1.0),
Z@1 = nth(Zp_arr, Channel, +0.0),
{S@1, Z@1}
end
)
end,
{ok, Result}
end
end.
-file("src/viva_tensor/quant/qat.gleam", 278).
?DOC(false).
-spec fake_quant_forward(
viva_tensor@tensor:tensor(),
quant_stats(),
quant_config()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
fake_quant_forward(Input, Stats, Config) ->
Data = viva_tensor@tensor:to_list(Input),
Shape = viva_tensor@tensor:shape(Input),
Scales = viva_tensor@tensor:to_list(erlang:element(2, Stats)),
Zps = viva_tensor@tensor:to_list(erlang:element(3, Stats)),
{Qmin, Qmax} = quant_range(Config),
gleam@result:'try'(
broadcast_per_element(Shape, Scales, Zps, Config),
fun(Scale_per_elem) ->
Pairs = gleam@list:zip(Data, Scale_per_elem),
Out = gleam@list:map(
Pairs,
fun(Pair) ->
{X, Sz} = Pair,
{Scale, Zp} = Sz,
Safe_scale = case Scale > +0.0 of
true ->
Scale;
false ->
1.0
end,
Q = erlang:round((case Safe_scale of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> X / Gleam@denominator
end) + Zp),
Q_clamped = gleam@int:max(
gleam@int:min(Q, float_to_int(Qmax)),
float_to_int(Qmin)
),
(erlang:float(Q_clamped) - Zp) * Safe_scale
end
),
{ok, {tensor, Out, Shape}}
end
).
-file("src/viva_tensor/quant/qat.gleam", 319).
?DOC(false).
-spec fake_quant_backward(
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
quant_stats(),
quant_config()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
fake_quant_backward(Grad_out, Input, Stats, Config) ->
Grad_data = viva_tensor@tensor:to_list(Grad_out),
In_data = viva_tensor@tensor:to_list(Input),
Shape = viva_tensor@tensor:shape(Input),
case erlang:length(Grad_data) =:= erlang:length(In_data) of
false ->
{error,
{invalid_shape,
<<"fake_quant_backward: grad/input shape mismatch"/utf8>>}};
true ->
Scales = viva_tensor@tensor:to_list(erlang:element(2, Stats)),
Zps = viva_tensor@tensor:to_list(erlang:element(3, Stats)),
{Qmin, Qmax} = quant_range(Config),
gleam@result:'try'(
broadcast_per_element(Shape, Scales, Zps, Config),
fun(Scale_per_elem) ->
Triples = gleam@list:zip(
Grad_data,
gleam@list:zip(In_data, Scale_per_elem)
),
Out = gleam@list:map(
Triples,
fun(T) ->
{G, Rest} = T,
{X, Sz} = Rest,
{Scale, Zp} = Sz,
Safe_scale = case Scale > +0.0 of
true ->
Scale;
false ->
1.0
end,
Q = (case Safe_scale of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> X / Gleam@denominator
end) + Zp,
case (Q >= Qmin) andalso (Q =< Qmax) of
true ->
G;
false ->
+0.0
end
end
),
{ok, {tensor, Out, Shape}}
end
)
end.
-file("src/viva_tensor/quant/qat.gleam", 429).
?DOC(false).
-spec compute_per_channel_scales(
viva_tensor@tensor:tensor(),
integer(),
integer()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
compute_per_channel_scales(Weight, Num_bits, Channel_axis) ->
Shape = viva_tensor@tensor:shape(Weight),
Rank = erlang:length(Shape),
case (Channel_axis < 0) orelse (Channel_axis >= Rank) of
true ->
{error,
{dimension_error,
<<"compute_per_channel_scales: channel_axis out of bounds"/utf8>>}};
false ->
Data = viva_tensor@tensor:to_list(Weight),
gleam@result:'try'(
split_along_axis(Data, Shape, Channel_axis),
fun(Channels) ->
Qmax = erlang:float(pow2(Num_bits - 1) - 1),
Scales = gleam@list:map(
Channels,
fun(Values) ->
Abs_max = begin
_pipe = Values,
_pipe@1 = gleam@list:map(
_pipe,
fun gleam@float:absolute_value/1
),
gleam@list:fold(
_pipe@1,
+0.0,
fun gleam@float:max/2
)
end,
case Abs_max > +0.0 of
true ->
case Qmax of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> Abs_max / Gleam@denominator
end;
false ->
1.0
end
end
),
{ok, {tensor, Scales, [erlang:length(Scales)]}}
end
)
end.
-file("src/viva_tensor/quant/qat.gleam", 485).
?DOC(false).
-spec qat_linear_init(integer(), integer(), integer(), integer()) -> qat_linear().
qat_linear_init(In_features, Out_features, Weight_bits, Activation_bits) ->
Weight = viva_tensor@tensor:zeros([Out_features, In_features]),
Bias = viva_tensor@tensor:zeros([Out_features]),
Weight_config = {quant_config, Weight_bits, true, true, 0},
Input_config = {quant_config, Activation_bits, true, false, 0},
Zero_scale = {tensor, gleam@list:repeat(1.0, Out_features), [Out_features]},
Zero_zp = {tensor, gleam@list:repeat(+0.0, Out_features), [Out_features]},
Weight_stats = {quant_stats, Zero_scale, Zero_zp},
{qat_linear,
Weight,
{some, Bias},
Weight_stats,
Weight_config,
none,
Input_config}.
-file("src/viva_tensor/quant/qat.gleam", 551).
?DOC(false).
-spec add_bias_row(viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
add_bias_row(Out, Bias) ->
Shape = viva_tensor@tensor:shape(Out),
case Shape of
[Batch, Features] ->
Bias_data = viva_tensor@tensor:to_list(Bias),
case erlang:length(Bias_data) =:= Features of
false ->
{error,
{invalid_shape,
<<"qat_linear: bias length != out_features"/utf8>>}};
true ->
Data = viva_tensor@tensor:to_list(Out),
Rows = gleam@list:sized_chunk(Data, Features),
Added = gleam@list:flat_map(
Rows,
fun(Row) ->
gleam@list:map2(
Row,
Bias_data,
fun(X, B) -> X + B end
)
end
),
{ok, {tensor, Added, [Batch, Features]}}
end;
_ ->
{error,
{dimension_error,
<<"qat_linear: expected [batch, out] output"/utf8>>}}
end.
-file("src/viva_tensor/quant/qat.gleam", 529).
?DOC(false).
-spec qat_linear_forward(qat_linear(), viva_tensor@tensor:tensor()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
qat_linear_forward(Layer, Input) ->
gleam@result:'try'(
fake_quant_forward(
erlang:element(2, Layer),
erlang:element(4, Layer),
erlang:element(5, Layer)
),
fun(Fq_weight) ->
Fq_input_result = case erlang:element(6, Layer) of
{some, Stats} ->
fake_quant_forward(Input, Stats, erlang:element(7, Layer));
none ->
{ok, Input}
end,
gleam@result:'try'(
Fq_input_result,
fun(Fq_input) ->
gleam@result:'try'(
viva_tensor@tensor:transpose(Fq_weight),
fun(W_t) ->
gleam@result:'try'(
viva_tensor@tensor:matmul(Fq_input, W_t),
fun(Out) -> case erlang:element(3, Layer) of
none ->
{ok, Out};
{some, B} ->
add_bias_row(Out, B)
end end
)
end
)
end
)
end
).
-file("src/viva_tensor/quant/qat.gleam", 576).
?DOC(false).
-spec qat_linear_calibrate(qat_linear(), viva_tensor@tensor:tensor()) -> {ok,
qat_linear()} |
{error, viva_tensor@core@error:tensor_error()}.
qat_linear_calibrate(Layer, Input) ->
gleam@result:'try'(
observe(Input, erlang:element(7, Layer)),
fun(Stats) ->
{ok,
{qat_linear,
erlang:element(2, Layer),
erlang:element(3, Layer),
erlang:element(4, Layer),
erlang:element(5, Layer),
{some, Stats},
erlang:element(7, Layer)}}
end
).