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@quant@awq.erl
-module(viva_tensor@quant@awq).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/quant/awq.gleam").
-export([default_config/0, collect_activation_stats/1, compute_awq_scales/2, apply_weight_transform/2, apply_activation_transform/2, quantize_awq/3, dequantize_awq/1, identify_salient_channels/2, benchmark_awq/0, main/0]).
-export_type([a_w_q_config/0, a_w_q_scales/0, a_w_q_tensor/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 a_w_q_config() :: {a_w_q_config, integer(), integer(), float(), boolean()}.
-type a_w_q_scales() :: {a_w_q_scales, list(float()), list(float()), float()}.
-type a_w_q_tensor() :: {a_w_q_tensor,
list(integer()),
a_w_q_scales(),
list(float()),
list(integer()),
list(integer()),
integer()}.
-file("src/viva_tensor/quant/awq.gleam", 89).
?DOC(false).
-spec default_config() -> a_w_q_config().
default_config() ->
{a_w_q_config, 4, 128, 0.5, false}.
-file("src/viva_tensor/quant/awq.gleam", 102).
?DOC(false).
-spec collect_activation_stats(list(list(float()))) -> list(float()).
collect_activation_stats(Activations_batch) ->
case Activations_batch of
[] ->
[];
[First | _] ->
Num_channels = erlang:length(First),
Initial = gleam@list:repeat(+0.0, Num_channels),
Sums = gleam@list:fold(
Activations_batch,
Initial,
fun(Acc, Activation) ->
gleam@list:map2(
Acc,
Activation,
fun(Sum, Act) ->
Sum + gleam@float:absolute_value(Act)
end
)
end
),
Num_samples = erlang:float(erlang:length(Activations_batch)),
gleam@list:map(Sums, fun(Sum@1) -> case Num_samples of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> Sum@1 / Gleam@denominator
end end)
end.
-file("src/viva_tensor/quant/awq.gleam", 497).
?DOC(false).
-spec float_power(float(), float()) -> float().
float_power(Base, Exp) ->
case gleam@float:power(Base, Exp) of
{ok, Result} ->
Result;
{error, _} ->
1.0
end.
-file("src/viva_tensor/quant/awq.gleam", 138).
?DOC(false).
-spec compute_awq_scales(list(float()), float()) -> a_w_q_scales().
compute_awq_scales(Activation_stats, Alpha) ->
Weight_scales = gleam@list:map(
Activation_stats,
fun(Stat) ->
Safe_stat = case Stat > +0.0 of
true ->
Stat;
false ->
1.0
end,
float_power(Safe_stat, Alpha)
end
),
{a_w_q_scales, Weight_scales, Activation_stats, Alpha}.
-file("src/viva_tensor/quant/awq.gleam", 161).
?DOC(false).
-spec apply_weight_transform(list(list(float())), a_w_q_scales()) -> list(list(float())).
apply_weight_transform(Weights, Scales) ->
gleam@list:map(
Weights,
fun(Row) ->
gleam@list:map2(
Row,
erlang:element(2, Scales),
fun(W, S) -> W * S end
)
end
).
-file("src/viva_tensor/quant/awq.gleam", 173).
?DOC(false).
-spec apply_activation_transform(list(float()), a_w_q_scales()) -> list(float()).
apply_activation_transform(Activations, Scales) ->
gleam@list:map2(
Activations,
erlang:element(2, Scales),
fun(X, S) -> case S > +0.0 of
true ->
case S of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> X / Gleam@denominator
end;
false ->
X
end end
).
-file("src/viva_tensor/quant/awq.gleam", 504).
?DOC(false).
-spec float_result_to_float({ok, float()} | {error, any()}, float()) -> float().
float_result_to_float(R, Default) ->
case R of
{ok, V} ->
V;
{error, _} ->
Default
end.
-file("src/viva_tensor/quant/awq.gleam", 252).
?DOC(false).
-spec symmetric_group_quantize(list(float()), integer(), integer()) -> {list(integer()),
list(float())}.
symmetric_group_quantize(Values, Bits, Group_size) ->
Qmax = begin
_pipe = gleam@float:power(2.0, erlang:float(Bits - 1)),
_pipe@1 = float_result_to_float(_pipe, 128.0),
(fun(X) -> X - 1.0 end)(_pipe@1)
end,
Groups = gleam@list:sized_chunk(Values, Group_size),
{Quantized_groups, Scales} = gleam@list:fold(
Groups,
{[], []},
fun(Acc, Group) ->
{Q_acc, S_acc} = Acc,
Max_abs = begin
_pipe@2 = Group,
_pipe@3 = gleam@list:map(
_pipe@2,
fun gleam@float:absolute_value/1
),
gleam@list:fold(_pipe@3, +0.0, fun gleam@float:max/2)
end,
Scale = case Max_abs > +0.0 of
true ->
case Max_abs of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> Qmax / Gleam@denominator
end;
false ->
1.0
end,
Quantized = gleam@list:map(
Group,
fun(V) ->
Scaled = V * Scale,
Clamped = gleam@float:clamp(Scaled, -1.0 * Qmax, Qmax),
erlang:round(Clamped)
end
),
{lists:append(Q_acc, Quantized), [Scale | S_acc]}
end
),
{Quantized_groups, lists:reverse(Scales)}.
-file("src/viva_tensor/quant/awq.gleam", 482).
?DOC(false).
-spec get_tensor_shape(viva_tensor@tensor:tensor()) -> list(integer()).
get_tensor_shape(T) ->
case T of
{tensor, _, Shape} ->
Shape;
{strided_tensor, _, Shape@1, _, _} ->
Shape@1;
{native_tensor, _, Shape@2} ->
Shape@2
end.
-file("src/viva_tensor/quant/awq.gleam", 198).
?DOC(false).
-spec quantize_awq(
viva_tensor@tensor:tensor(),
list(list(float())),
a_w_q_config()
) -> a_w_q_tensor().
quantize_awq(Weights, Calibration_data, Config) ->
Weight_data = viva_tensor@tensor:to_list(Weights),
Shape = get_tensor_shape(Weights),
{_, In_features} = case Shape of
[O, I] ->
{O, I};
_ ->
{1, erlang:length(Weight_data)}
end,
Weight_matrix = gleam@list:sized_chunk(Weight_data, In_features),
Activation_stats = collect_activation_stats(Calibration_data),
Awq_scales = compute_awq_scales(Activation_stats, erlang:element(4, Config)),
Transformed_weights = apply_weight_transform(Weight_matrix, Awq_scales),
Flat_transformed = lists:append(Transformed_weights),
{Quantized, Quant_scales} = symmetric_group_quantize(
Flat_transformed,
erlang:element(2, Config),
erlang:element(3, Config)
),
Num_elements = erlang:length(Flat_transformed),
Num_groups = case erlang:element(3, Config) of
0 -> 0;
Gleam@denominator -> ((Num_elements + erlang:element(3, Config)) - 1)
div Gleam@denominator
end,
Data_bytes = ((Num_elements * erlang:element(2, Config)) + 7) div 8,
Scale_bytes = Num_groups * 2,
Awq_scale_bytes = In_features * 2,
Memory = (Data_bytes + Scale_bytes) + Awq_scale_bytes,
{a_w_q_tensor, Quantized, Awq_scales, Quant_scales, [], Shape, Memory}.
-file("src/viva_tensor/quant/awq.gleam", 490).
?DOC(false).
-spec get_at_index_float(list(float()), integer(), float()) -> float().
get_at_index_float(Lst, Idx, Default) ->
case gleam@list:drop(Lst, Idx) of
[First | _] ->
First;
[] ->
Default
end.
-file("src/viva_tensor/quant/awq.gleam", 297).
?DOC(false).
-spec dequantize_awq(a_w_q_tensor()) -> viva_tensor@tensor:tensor().
dequantize_awq(Awq) ->
Group_size = case erlang:element(4, Awq) of
[] ->
erlang:length(erlang:element(2, Awq));
_ ->
case erlang:length(erlang:element(4, Awq)) of
0 -> 0;
Gleam@denominator -> erlang:length(erlang:element(2, Awq)) div Gleam@denominator
end
end,
Groups = gleam@list:sized_chunk(erlang:element(2, Awq), Group_size),
Dequantized = begin
_pipe = gleam@list:index_map(
Groups,
fun(Group, Idx) ->
Scale = get_at_index_float(erlang:element(4, Awq), Idx, 1.0),
gleam@list:map(Group, fun(Q) -> case Scale of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> erlang:float(Q) / Gleam@denominator@1
end end)
end
),
lists:append(_pipe)
end,
In_features = case erlang:element(6, Awq) of
[_, I] ->
I;
_ ->
1
end,
Weight_matrix = gleam@list:sized_chunk(Dequantized, In_features),
Restored = begin
_pipe@1 = gleam@list:map(
Weight_matrix,
fun(Row) ->
gleam@list:map2(
Row,
erlang:element(2, erlang:element(3, Awq)),
fun(W, S) -> case S > +0.0 of
true ->
case S of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@2 -> W / Gleam@denominator@2
end;
false ->
W
end end
)
end
),
lists:append(_pipe@1)
end,
{tensor, Restored, erlang:element(6, Awq)}.
-file("src/viva_tensor/quant/awq.gleam", 342).
?DOC(false).
-spec identify_salient_channels(list(float()), float()) -> list(integer()).
identify_salient_channels(Activation_stats, Top_percent) ->
N = erlang:length(Activation_stats),
K = begin
_pipe = erlang:round((erlang:float(N) * Top_percent) / 100.0),
gleam@int:max(_pipe, 1)
end,
_pipe@1 = Activation_stats,
_pipe@2 = gleam@list:index_map(_pipe@1, fun(Stat, Idx) -> {Idx, Stat} end),
_pipe@3 = gleam@list:sort(
_pipe@2,
fun(A, B) ->
gleam@float:compare(erlang:element(2, B), erlang:element(2, A))
end
),
_pipe@4 = gleam@list:take(_pipe@3, K),
gleam@list:map(_pipe@4, fun(Pair) -> erlang:element(1, Pair) end).
-file("src/viva_tensor/quant/awq.gleam", 511).
?DOC(false).
-spec float_to_string(float()) -> binary().
float_to_string(F) ->
Rounded = erlang:float(erlang:round(F * 10000.0)) / 10000.0,
gleam_stdlib:float_to_string(Rounded).
-file("src/viva_tensor/quant/awq.gleam", 523).
?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/awq.gleam", 519).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/quant/awq.gleam", 365).
?DOC(false).
-spec benchmark_awq() -> nil.
benchmark_awq() ->
gleam_stdlib:println(
<<"====================================================================="/utf8>>
),
gleam_stdlib:println(<<" AWQ - Lin et al. (2024) MLSys Best Paper"/utf8>>),
gleam_stdlib:println(
<<" The key insight: 1% of weights matter 10x more"/utf8>>
),
gleam_stdlib:println(
<<"=====================================================================\n"/utf8>>
),
gleam_stdlib:println(<<"--- The Algorithm ---"/utf8>>),
gleam_stdlib:println(
<<" 1. Collect activation statistics (calibration)"/utf8>>
),
gleam_stdlib:println(
<<" 2. Identify salient channels (high activation = important)"/utf8>>
),
gleam_stdlib:println(
<<" 3. Scale salient weights UP before quantizing"/utf8>>
),
gleam_stdlib:println(
<<" 4. Scale activations DOWN at runtime (mathematically equivalent)"/utf8>>
),
gleam_stdlib:println(
<<" Result: Protected channels get more quantization precision"/utf8>>
),
gleam_stdlib:println(<<""/utf8>>),
Weights = viva_tensor@tensor:random_uniform([512, 256]),
Calibration_data = begin
_pipe = range_int(1, 100),
gleam@list:map(
_pipe,
fun(_) -> _pipe@1 = viva_tensor@tensor:random_uniform([256]),
viva_tensor@tensor:to_list(_pipe@1) end
)
end,
Config = default_config(),
gleam_stdlib:println(<<"--- Calibration ---"/utf8>>),
Activation_stats = collect_activation_stats(Calibration_data),
gleam_stdlib:println(<<" Samples: 100"/utf8>>),
gleam_stdlib:println(<<" Features: 256"/utf8>>),
Salient_channels = identify_salient_channels(Activation_stats, 1.0),
gleam_stdlib:println(
<<" Salient channels (top 1%): "/utf8,
(erlang:integer_to_binary(erlang:length(Salient_channels)))/binary>>
),
gleam_stdlib:println(<<" Top 5 most salient:"/utf8>>),
_pipe@2 = Salient_channels,
_pipe@3 = gleam@list:take(_pipe@2, 5),
gleam@list:each(
_pipe@3,
fun(Idx) ->
Stat = get_at_index_float(Activation_stats, Idx, +0.0),
gleam_stdlib:println(
<<<<<<" Channel "/utf8,
(erlang:integer_to_binary(Idx))/binary>>/binary,
": "/utf8>>/binary,
(float_to_string(Stat))/binary>>
)
end
),
gleam_stdlib:println(<<"\n--- AWQ Quantization ---"/utf8>>),
{Time_awq, Awq_tensor} = timer:tc(
fun() -> quantize_awq(Weights, Calibration_data, Config) end
),
Original_bytes = (512 * 256) * 4,
Ratio = case erlang:float(erlang:element(7, Awq_tensor)) of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> erlang:float(Original_bytes) / Gleam@denominator
end,
gleam_stdlib:println(
<<<<" Time: "/utf8,
(erlang:integer_to_binary(Time_awq div 1000))/binary>>/binary,
"ms"/utf8>>
),
gleam_stdlib:println(
<<<<" Original: "/utf8,
(erlang:integer_to_binary(Original_bytes div 1024))/binary>>/binary,
" KB"/utf8>>
),
gleam_stdlib:println(
<<<<" Compressed: "/utf8,
(erlang:integer_to_binary(
erlang:element(7, Awq_tensor) div 1024
))/binary>>/binary,
" KB"/utf8>>
),
gleam_stdlib:println(
<<<<" Compression: "/utf8, (float_to_string(Ratio))/binary>>/binary,
"x"/utf8>>
),
gleam_stdlib:println(<<"\n--- Error Analysis ---"/utf8>>),
Decompressed = dequantize_awq(Awq_tensor),
Orig_data = viva_tensor@tensor:to_list(Weights),
Decomp_data = viva_tensor@tensor:to_list(Decompressed),
Errors = gleam@list:map2(
Orig_data,
Decomp_data,
fun(O, D) -> gleam@float:absolute_value(O - D) end
),
Mean_error = case Errors of
[] ->
+0.0;
_ ->
case erlang:float(erlang:length(Errors)) of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> gleam@list:fold(
Errors,
+0.0,
fun gleam@float:add/2
)
/ Gleam@denominator@1
end
end,
Max_error = gleam@list:fold(Errors, +0.0, fun gleam@float:max/2),
gleam_stdlib:println(
<<" Mean error: "/utf8, (float_to_string(Mean_error))/binary>>
),
gleam_stdlib:println(
<<" Max error: "/utf8, (float_to_string(Max_error))/binary>>
),
gleam_stdlib:println(
<<"\n--- Why AWQ Beats Standard Quantization ---"/utf8>>
),
gleam_stdlib:println(<<" Standard: All channels quantized equally"/utf8>>),
gleam_stdlib:println(<<" AWQ: Salient channels get more precision"/utf8>>),
gleam_stdlib:println(
<<" Same compression ratio, MUCH lower perplexity"/utf8>>
),
gleam_stdlib:println(
<<"\n====================================================================="/utf8>>
),
gleam_stdlib:println(<<" AWQ IN PRODUCTION"/utf8>>),
gleam_stdlib:println(<<""/utf8>>),
gleam_stdlib:println(<<" LLaMA-7B:"/utf8>>),
gleam_stdlib:println(<<" - FP16: 14GB"/utf8>>),
gleam_stdlib:println(<<" - AWQ-4bit: 3.5GB (fits on RTX 3060!)"/utf8>>),
gleam_stdlib:println(<<" - Perplexity loss: <0.5%"/utf8>>),
gleam_stdlib:println(<<""/utf8>>),
gleam_stdlib:println(<<" LLaMA-70B:"/utf8>>),
gleam_stdlib:println(<<" - FP16: 140GB (needs 8x A100)"/utf8>>),
gleam_stdlib:println(<<" - AWQ-4bit: 35GB (fits on single A100!)"/utf8>>),
gleam_stdlib:println(<<" - Perplexity loss: <1%"/utf8>>),
gleam_stdlib:println(<<""/utf8>>),
gleam_stdlib:println(
<<" Zero runtime overhead - transform is pre-computed."/utf8>>
),
gleam_stdlib:println(
<<"====================================================================="/utf8>>
).
-file("src/viva_tensor/quant/awq.gleam", 361).
?DOC(false).
-spec main() -> nil.
main() ->
benchmark_awq().