-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 ).