-module(viva_tensor@metrics@classification). -compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]). -define(FILEPATH, "src/viva_tensor/metrics/classification.gleam"). -export([accuracy/2, confusion_matrix/3, precision/4, recall/4, f1/4, top_k_accuracy/3, iou_per_class/3, mean_iou/3]). -export_type([average/0, class_stats/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 average() :: micro | macro | weighted. -type class_stats() :: {class_stats, list(integer()), list(integer()), list(integer()), list(integer())}. -file("src/viva_tensor/metrics/classification.gleam", 284). ?DOC(false). -spec to_indices(list(float())) -> {ok, list(integer())} | {error, viva_tensor@core@error:tensor_error()}. to_indices(Xs) -> gleam@list:try_map( Xs, fun(X) -> I = erlang:round(X), case I >= 0 of true -> {ok, I}; false -> {error, {index_out_of_bounds, I, 0}} end end ). -file("src/viva_tensor/metrics/classification.gleam", 256). ?DOC(false). -spec pair_indices(viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor()) -> {ok, {list(integer()), list(integer())}} | {error, viva_tensor@core@error:tensor_error()}. pair_indices(Predictions, Targets) -> Pred_shape = viva_tensor@tensor:shape(Predictions), Target_shape = viva_tensor@tensor:shape(Targets), _pipe = case Pred_shape =:= Target_shape of true -> {ok, nil}; false -> {error, {shape_mismatch, Target_shape, Pred_shape}} end, _pipe@1 = gleam@result:'try'(_pipe, fun(_) -> case Pred_shape of [_] -> {ok, nil}; _ -> {error, {invalid_shape, <<"expected 1D tensors, got "/utf8, (viva_tensor@core@error:shape_to_string( Pred_shape ))/binary>>}} end end), gleam@result:'try'( _pipe@1, fun(_) -> gleam@result:'try'( viva_tensor@tensor:try_to_list(Predictions), fun(Pred_data) -> gleam@result:'try'( viva_tensor@tensor:try_to_list(Targets), fun(Target_data) -> gleam@result:'try'( to_indices(Pred_data), fun(Preds) -> gleam@result:'try'( to_indices(Target_data), fun(Tgts) -> {ok, {Preds, Tgts}} end ) end ) end ) end ) end ). -file("src/viva_tensor/metrics/classification.gleam", 41). ?DOC(false). -spec accuracy(viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor()) -> {ok, float()} | {error, viva_tensor@core@error:tensor_error()}. accuracy(Predictions, Targets) -> gleam@result:'try'( pair_indices(Predictions, Targets), fun(_use0) -> {Preds, Tgts} = _use0, N = erlang:length(Preds), case N of 0 -> {error, {invalid_shape, <<"accuracy: empty inputs"/utf8>>}}; _ -> Matches = begin _pipe = gleam@list:zip(Preds, Tgts), gleam@list:fold( _pipe, 0, fun(Acc, Pair) -> {P, T} = Pair, case P =:= T of true -> Acc + 1; false -> Acc end end ) end, {ok, case erlang:float(N) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> erlang:float(Matches) / Gleam@denominator end} end end ). -file("src/viva_tensor/metrics/classification.gleam", 317). ?DOC(false). -spec increment_at(list(integer()), integer()) -> list(integer()). increment_at(Xs, Idx) -> gleam@list:index_map(Xs, fun(Value, I) -> case I =:= Idx of true -> Value + 1; false -> Value end end). -file("src/viva_tensor/metrics/classification.gleam", 294). ?DOC(false). -spec build_counts(list(integer()), list(integer()), integer()) -> {ok, list(float())} | {error, viva_tensor@core@error:tensor_error()}. build_counts(Preds, Tgts, Num_classes) -> Size = Num_classes * Num_classes, Zeros = gleam@list:repeat(0, Size), Pairs = gleam@list:zip(Tgts, Preds), _pipe = gleam@list:try_fold( Pairs, Zeros, fun(Acc, Pair) -> {T, P} = Pair, case T >= Num_classes of true -> {error, {index_out_of_bounds, T, Num_classes}}; false -> case P >= Num_classes of true -> {error, {index_out_of_bounds, P, Num_classes}}; false -> {ok, increment_at(Acc, (T * Num_classes) + P)} end end end ), gleam@result:map( _pipe, fun(Ints) -> gleam@list:map(Ints, fun erlang:float/1) end ). -file("src/viva_tensor/metrics/classification.gleam", 72). ?DOC(false). -spec confusion_matrix( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer() ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. confusion_matrix(Predictions, Targets, Num_classes) -> gleam@result:'try'(case Num_classes > 0 of true -> {ok, nil}; false -> {error, {invalid_shape, <<"confusion_matrix: num_classes must be > 0"/utf8>>}} end, fun(_) -> gleam@result:'try'( pair_indices(Predictions, Targets), fun(_use0) -> {Preds, Tgts} = _use0, gleam@result:'try'( build_counts(Preds, Tgts, Num_classes), fun(Counts) -> viva_tensor@tensor:matrix( Num_classes, Num_classes, Counts ) end ) end ) end). -file("src/viva_tensor/metrics/classification.gleam", 425). ?DOC(false). -spec int_sum(list(integer())) -> integer(). int_sum(Xs) -> gleam@list:fold(Xs, 0, fun(Acc, V) -> Acc + V end). -file("src/viva_tensor/metrics/classification.gleam", 437). ?DOC(false). -spec weighted_mean(list(float()), list(integer())) -> float(). weighted_mean(Values, Weights) -> Total = int_sum(Weights), case Total of 0 -> +0.0; _ -> Weighted = begin _pipe = gleam@list:zip(Values, Weights), gleam@list:fold( _pipe, +0.0, fun(Acc, Pair) -> {V, W} = Pair, Acc + (V * erlang:float(W)) end ) end, case erlang:float(Total) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> Weighted / Gleam@denominator end end. -file("src/viva_tensor/metrics/classification.gleam", 418). ?DOC(false). -spec safe_ratio(float(), float()) -> float(). safe_ratio(Num, Denom) -> case Denom > +0.0 of true -> case Denom of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> Num / Gleam@denominator end; false -> +0.0 end. -file("src/viva_tensor/metrics/classification.gleam", 397). ?DOC(false). -spec per_class_ratio(list(integer()), list(integer())) -> list(float()). per_class_ratio(Tp, Other) -> _pipe = gleam@list:zip(Tp, Other), gleam@list:map( _pipe, fun(Pair) -> {Tp_c, Other_c} = Pair, safe_ratio(erlang:float(Tp_c), erlang:float(Tp_c + Other_c)) end ). -file("src/viva_tensor/metrics/classification.gleam", 429). ?DOC(false). -spec mean_floats(list(float())) -> float(). mean_floats(Xs) -> N = erlang:length(Xs), case N of 0 -> +0.0; _ -> case erlang:float(N) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> gleam@list:fold( Xs, +0.0, fun(Acc, V) -> Acc + V end ) / Gleam@denominator end end. -file("src/viva_tensor/metrics/classification.gleam", 374). ?DOC(false). -spec aggregate(list(integer()), list(integer()), list(integer()), average()) -> float(). aggregate(Tp, Other, Support, Average) -> case Average of micro -> Sum_tp = int_sum(Tp), Sum_other = int_sum(Other), safe_ratio(erlang:float(Sum_tp), erlang:float(Sum_tp + Sum_other)); macro -> Per_class = per_class_ratio(Tp, Other), mean_floats(Per_class); weighted -> Per_class@1 = per_class_ratio(Tp, Other), weighted_mean(Per_class@1, Support) end. -file("src/viva_tensor/metrics/classification.gleam", 326). ?DOC(false). -spec class_stats( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer() ) -> {ok, class_stats()} | {error, viva_tensor@core@error:tensor_error()}. class_stats(Predictions, Targets, Num_classes) -> _pipe = case Num_classes > 0 of true -> {ok, nil}; false -> {error, {invalid_shape, <<"num_classes must be > 0"/utf8>>}} end, _pipe@1 = gleam@result:'try'( _pipe, fun(_) -> pair_indices(Predictions, Targets) end ), gleam@result:'try'( _pipe@1, fun(Pair) -> {Preds, Tgts} = Pair, Zeros = gleam@list:repeat(0, Num_classes), Folded = gleam@list:try_fold( gleam@list:zip(Tgts, Preds), {Zeros, Zeros, Zeros, Zeros}, fun(Acc, Pair@1) -> {T, P} = Pair@1, {Tp, Fp, Fn_, Support} = Acc, case T >= Num_classes of true -> {error, {index_out_of_bounds, T, Num_classes}}; false -> case P >= Num_classes of true -> {error, {index_out_of_bounds, P, Num_classes}}; false -> Support2 = increment_at(Support, T), case T =:= P of true -> {ok, {increment_at(Tp, T), Fp, Fn_, Support2}}; false -> {ok, {Tp, increment_at(Fp, P), increment_at(Fn_, T), Support2}} end end end end ), gleam@result:'try'( Folded, fun(Stats) -> {Tp@1, Fp@1, Fn_@1, Support@1} = Stats, {ok, {class_stats, Tp@1, Fp@1, Fn_@1, Support@1}} end ) end ). -file("src/viva_tensor/metrics/classification.gleam", 94). ?DOC(false). -spec precision( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer(), average() ) -> {ok, float()} | {error, viva_tensor@core@error:tensor_error()}. precision(Predictions, Targets, Num_classes, Average) -> gleam@result:'try'( class_stats(Predictions, Targets, Num_classes), fun(Stats) -> {class_stats, Tp, Fp, _, Support} = Stats, {ok, aggregate(Tp, Fp, Support, Average)} end ). -file("src/viva_tensor/metrics/classification.gleam", 108). ?DOC(false). -spec recall( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer(), average() ) -> {ok, float()} | {error, viva_tensor@core@error:tensor_error()}. recall(Predictions, Targets, Num_classes, Average) -> gleam@result:'try'( class_stats(Predictions, Targets, Num_classes), fun(Stats) -> {class_stats, Tp, _, Fn_, Support} = Stats, {ok, aggregate(Tp, Fn_, Support, Average)} end ). -file("src/viva_tensor/metrics/classification.gleam", 405). ?DOC(false). -spec per_class_f1(list(integer()), list(integer()), list(integer())) -> list(float()). per_class_f1(Tp, Fp, Fn_) -> _pipe = gleam@list:zip(Tp, gleam@list:zip(Fp, Fn_)), gleam@list:map( _pipe, fun(Triple) -> {Tp_c, {Fp_c, Fn_c}} = Triple, P = safe_ratio(erlang:float(Tp_c), erlang:float(Tp_c + Fp_c)), R = safe_ratio(erlang:float(Tp_c), erlang:float(Tp_c + Fn_c)), case (P + R) > +0.0 of true -> case (P + R) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> (2.0 * P) * R / Gleam@denominator end; false -> +0.0 end end ). -file("src/viva_tensor/metrics/classification.gleam", 124). ?DOC(false). -spec f1( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer(), average() ) -> {ok, float()} | {error, viva_tensor@core@error:tensor_error()}. f1(Predictions, Targets, Num_classes, Average) -> gleam@result:'try'( class_stats(Predictions, Targets, Num_classes), fun(Stats) -> {class_stats, Tp, Fp, Fn_, Support} = Stats, case Average of micro -> Sum_tp = int_sum(Tp), Sum_fp = int_sum(Fp), Sum_fn = int_sum(Fn_), {ok, safe_ratio( erlang:float(Sum_tp), erlang:float(Sum_tp + ((Sum_fp + Sum_fn) div 2)) )}; macro -> Per_class = per_class_f1(Tp, Fp, Fn_), {ok, mean_floats(Per_class)}; weighted -> Per_class@1 = per_class_f1(Tp, Fp, Fn_), {ok, weighted_mean(Per_class@1, Support)} end end ). -file("src/viva_tensor/metrics/classification.gleam", 466). ?DOC(false). -spec row_topk_contains(list(float()), integer(), integer()) -> boolean(). row_topk_contains(Row, K, Target) -> Indexed = gleam@list:index_map(Row, fun(Value, Idx) -> {Value, Idx} end), Sorted = gleam@list:sort( Indexed, fun(A, B) -> {Va, Ia} = A, {Vb, Ib} = B, case gleam@float:compare(Vb, Va) of eq -> gleam@int:compare(Ia, Ib); Ord -> Ord end end ), _pipe = Sorted, _pipe@1 = gleam@list:take(_pipe, K), gleam@list:any( _pipe@1, fun(Pair) -> {_, Idx@1} = Pair, Idx@1 =:= Target end ). -file("src/viva_tensor/metrics/classification.gleam", 453). ?DOC(false). -spec chunk_rows(list(float()), integer()) -> list(list(float())). chunk_rows(Data, Cols) -> case Data of [] -> []; _ -> Row = gleam@list:take(Data, Cols), Rest = gleam@list:drop(Data, Cols), [Row | chunk_rows(Rest, Cols)] end. -file("src/viva_tensor/metrics/classification.gleam", 160). ?DOC(false). -spec top_k_accuracy( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer() ) -> {ok, float()} | {error, viva_tensor@core@error:tensor_error()}. top_k_accuracy(Logits, Targets, K) -> _pipe = case K > 0 of true -> {ok, nil}; false -> {error, {invalid_shape, <<"top_k_accuracy: k must be > 0"/utf8>>}} end, _pipe@1 = gleam@result:'try'( _pipe, fun(_) -> case viva_tensor@tensor:shape(Logits) of [Batch, Num_classes] -> {ok, {Batch, Num_classes}}; Other -> {error, {invalid_shape, <<"top_k_accuracy: logits must be 2D, got "/utf8, (viva_tensor@core@error:shape_to_string(Other))/binary>>}} end end ), _pipe@2 = gleam@result:'try'( _pipe@1, fun(Dims) -> {Batch@1, Num_classes@1} = Dims, case viva_tensor@tensor:shape(Targets) of [N] when N =:= Batch@1 -> {ok, {Batch@1, Num_classes@1}}; Other@1 -> {error, {shape_mismatch, [Batch@1], Other@1}} end end ), gleam@result:'try'( _pipe@2, fun(Dims@1) -> {Batch@2, Num_classes@2} = Dims@1, gleam@result:'try'( viva_tensor@tensor:try_to_list(Logits), fun(Logit_data) -> gleam@result:'try'( viva_tensor@tensor:try_to_list(Targets), fun(Target_data) -> gleam@result:'try'( to_indices(Target_data), fun(Target_idx) -> Effective_k = case K > Num_classes@2 of true -> Num_classes@2; false -> K end, Rows = chunk_rows(Logit_data, Num_classes@2), Hits = begin _pipe@3 = gleam@list:zip( Rows, Target_idx ), gleam@list:fold( _pipe@3, 0, fun(Acc, Pair) -> {Row, T} = Pair, case row_topk_contains( Row, Effective_k, T ) of true -> Acc + 1; false -> Acc end end ) end, case Batch@2 of 0 -> {error, {invalid_shape, <<"top_k_accuracy: empty batch"/utf8>>}}; _ -> {ok, case erlang:float(Batch@2) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> erlang:float( Hits ) / Gleam@denominator end} end end ) end ) end ) end ). -file("src/viva_tensor/metrics/classification.gleam", 217). ?DOC(false). -spec iou_per_class( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer() ) -> {ok, list(float())} | {error, viva_tensor@core@error:tensor_error()}. iou_per_class(Predictions, Targets, Num_classes) -> gleam@result:'try'( class_stats(Predictions, Targets, Num_classes), fun(Stats) -> {class_stats, Tp, Fp, Fn_, _} = Stats, Triples = gleam@list:zip(Tp, gleam@list:zip(Fp, Fn_)), Ious = gleam@list:map( Triples, fun(Triple) -> {Tp_c, {Fp_c, Fn_c}} = Triple, Denom = (Tp_c + Fp_c) + Fn_c, case Denom of 0 -> +0.0; _ -> case erlang:float(Denom) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> erlang:float(Tp_c) / Gleam@denominator end end end ), {ok, Ious} end ). -file("src/viva_tensor/metrics/classification.gleam", 240). ?DOC(false). -spec mean_iou( viva_tensor@tensor:tensor(), viva_tensor@tensor:tensor(), integer() ) -> {ok, float()} | {error, viva_tensor@core@error:tensor_error()}. mean_iou(Predictions, Targets, Num_classes) -> gleam@result:'try'( iou_per_class(Predictions, Targets, Num_classes), fun(Ious) -> {ok, mean_floats(Ious)} end ).