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