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@metrics@classification.erl
-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
).