Current section

Files

Jump to
viva_tensor src viva_tensor@data@dataloader.erl
Raw

src/viva_tensor@data@dataloader.erl

-module(viva_tensor@data@dataloader).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/data/dataloader.gleam").
-export([dataset_from_samples/1, dataset_from_lists/2, dataset_len/1, dataset_get/2, data_loader_new/4, data_loader_batches/1, data_loader_len/1]).
-export_type([sample/0, dataset/0, batch/0, data_loader/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 sample() :: {sample,
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()}.
-opaque dataset() :: {dataset, list(sample())}.
-type batch() :: {batch,
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor()}.
-type data_loader() :: {data_loader, dataset(), integer(), boolean(), boolean()}.
-file("src/viva_tensor/data/dataloader.gleam", 75).
?DOC(false).
-spec dataset_from_samples(list(sample())) -> dataset().
dataset_from_samples(Samples) ->
{dataset, Samples}.
-file("src/viva_tensor/data/dataloader.gleam", 256).
?DOC(false).
-spec validate_uniform_shapes(list(viva_tensor@tensor:tensor()), binary()) -> {ok,
nil} |
{error, viva_tensor@core@error:tensor_error()}.
validate_uniform_shapes(Tensors, _) ->
case Tensors of
[] ->
{ok, nil};
[First | Rest] ->
Expected = viva_tensor@tensor:shape(First),
case gleam@list:find(
Rest,
fun(T) -> viva_tensor@tensor:shape(T) /= Expected end
) of
{ok, Bad} ->
{error,
{shape_mismatch,
Expected,
viva_tensor@tensor:shape(Bad)}};
{error, _} ->
{ok, nil}
end
end.
-file("src/viva_tensor/data/dataloader.gleam", 92).
?DOC(false).
-spec dataset_from_lists(
list(viva_tensor@tensor:tensor()),
list(viva_tensor@tensor:tensor())
) -> {ok, dataset()} | {error, viva_tensor@core@error:tensor_error()}.
dataset_from_lists(Inputs, Targets) ->
N_inputs = erlang:length(Inputs),
N_targets = erlang:length(Targets),
case N_inputs =:= N_targets of
false ->
{error, {shape_mismatch, [N_inputs], [N_targets]}};
true ->
gleam@result:'try'(
validate_uniform_shapes(Inputs, <<"input"/utf8>>),
fun(_) ->
gleam@result:'try'(
validate_uniform_shapes(Targets, <<"target"/utf8>>),
fun(_) ->
Samples = gleam@list:map2(
Inputs,
Targets,
fun(X, Y) -> {sample, X, Y} end
),
{ok, {dataset, Samples}}
end
)
end
)
end.
-file("src/viva_tensor/data/dataloader.gleam", 118).
?DOC(false).
-spec dataset_len(dataset()) -> integer().
dataset_len(D) ->
erlang:length(erlang:element(2, D)).
-file("src/viva_tensor/data/dataloader.gleam", 248).
?DOC(false).
-spec list_at(list(VQO), integer()) -> {ok, VQO} | {error, nil}.
list_at(Xs, Index) ->
case {Xs, Index} of
{[], _} ->
{error, nil};
{[Head | _], 0} ->
{ok, Head};
{[_ | Rest], _} ->
list_at(Rest, Index - 1)
end.
-file("src/viva_tensor/data/dataloader.gleam", 136).
?DOC(false).
-spec dataset_get(dataset(), integer()) -> {ok, sample()} |
{error, viva_tensor@core@error:tensor_error()}.
dataset_get(D, Index) ->
N = erlang:length(erlang:element(2, D)),
case N of
0 ->
{error, {index_out_of_bounds, Index, 0}};
_ ->
Resolved = case Index < 0 of
true ->
Index + N;
false ->
Index
end,
case (Resolved >= 0) andalso (Resolved < N) of
false ->
{error, {index_out_of_bounds, Index, N}};
true ->
case list_at(erlang:element(2, D), Resolved) of
{ok, Sample} ->
{ok, Sample};
{error, _} ->
{error, {index_out_of_bounds, Index, N}}
end
end
end.
-file("src/viva_tensor/data/dataloader.gleam", 167).
?DOC(false).
-spec data_loader_new(dataset(), integer(), boolean(), boolean()) -> data_loader().
data_loader_new(Dataset, Batch_size, Shuffle, Drop_last) ->
{data_loader, Dataset, Batch_size, Shuffle, Drop_last}.
-file("src/viva_tensor/data/dataloader.gleam", 316).
?DOC(false).
-spec stack_tensors(list(viva_tensor@tensor:tensor()), list(integer())) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
stack_tensors(Tensors, Expected_shape) ->
case gleam@list:find(
Tensors,
fun(T) -> viva_tensor@tensor:shape(T) /= Expected_shape end
) of
{ok, Bad} ->
{error,
{shape_mismatch, Expected_shape, viva_tensor@tensor:shape(Bad)}};
{error, _} ->
Batch = erlang:length(Tensors),
Flat = begin
_pipe = Tensors,
_pipe@1 = gleam@list:map(
_pipe,
fun viva_tensor@tensor:to_list/1
),
lists:append(_pipe@1)
end,
Shape = [Batch | Expected_shape],
{ok, {tensor, Flat, Shape}}
end.
-file("src/viva_tensor/data/dataloader.gleam", 297).
?DOC(false).
-spec stack_samples(list(sample())) -> {ok, batch()} |
{error, viva_tensor@core@error:tensor_error()}.
stack_samples(Samples) ->
case Samples of
[] ->
{error, {invalid_shape, <<"cannot stack an empty batch"/utf8>>}};
[First | _] ->
Input_shape = viva_tensor@tensor:shape(erlang:element(2, First)),
Target_shape = viva_tensor@tensor:shape(erlang:element(3, First)),
gleam@result:'try'(
stack_tensors(
gleam@list:map(Samples, fun(S) -> erlang:element(2, S) end),
Input_shape
),
fun(Inputs) ->
gleam@result:'try'(
stack_tensors(
gleam@list:map(
Samples,
fun(S@1) -> erlang:element(3, S@1) end
),
Target_shape
),
fun(Targets) -> {ok, {batch, Inputs, Targets}} end
)
end
)
end.
-file("src/viva_tensor/data/dataloader.gleam", 284).
?DOC(false).
-spec stack_groups(list(list(sample())), list(batch())) -> {ok, list(batch())} |
{error, viva_tensor@core@error:tensor_error()}.
stack_groups(Groups, Acc) ->
case Groups of
[] ->
{ok, lists:reverse(Acc)};
[Group | Rest] ->
gleam@result:'try'(
stack_samples(Group),
fun(Batch) -> stack_groups(Rest, [Batch | Acc]) end
)
end.
-file("src/viva_tensor/data/dataloader.gleam", 273).
?DOC(false).
-spec chunk(list(VQV), integer()) -> list(list(VQV)).
chunk(Xs, Size) ->
case Xs of
[] ->
[];
_ ->
Head = gleam@list:take(Xs, Size),
Tail = gleam@list:drop(Xs, Size),
[Head | chunk(Tail, Size)]
end.
-file("src/viva_tensor/data/dataloader.gleam", 341).
?DOC(false).
-spec shuffle_samples(list(sample())) -> list(sample()).
shuffle_samples(Samples) ->
_pipe = Samples,
_pipe@1 = gleam@list:map(
_pipe,
fun(S) -> {gleam@int:random(1000000000), S} end
),
_pipe@2 = gleam@list:sort(
_pipe@1,
fun(A, B) ->
{Ka, _} = A,
{Kb, _} = B,
case Ka < Kb of
true ->
lt;
false ->
case Ka > Kb of
true ->
gt;
false ->
eq
end
end
end
),
gleam@list:map(
_pipe@2,
fun(Pair) ->
{_, S@1} = Pair,
S@1
end
).
-file("src/viva_tensor/data/dataloader.gleam", 200).
?DOC(false).
-spec data_loader_batches(data_loader()) -> {ok, list(batch())} |
{error, viva_tensor@core@error:tensor_error()}.
data_loader_batches(Loader) ->
case erlang:element(3, Loader) =< 0 of
true ->
{error, {invalid_shape, <<"batch_size must be > 0"/utf8>>}};
false ->
Samples = erlang:element(2, erlang:element(2, Loader)),
Ordered = case erlang:element(4, Loader) of
true ->
shuffle_samples(Samples);
false ->
Samples
end,
Groups = chunk(Ordered, erlang:element(3, Loader)),
Kept = case erlang:element(5, Loader) of
true ->
gleam@list:filter(
Groups,
fun(G) ->
erlang:length(G) =:= erlang:element(3, Loader)
end
);
false ->
Groups
end,
stack_groups(Kept, [])
end.
-file("src/viva_tensor/data/dataloader.gleam", 231).
?DOC(false).
-spec data_loader_len(data_loader()) -> integer().
data_loader_len(Loader) ->
case erlang:element(3, Loader) =< 0 of
true ->
0;
false ->
N = erlang:length(erlang:element(2, erlang:element(2, Loader))),
Full = case erlang:element(3, Loader) of
0 -> 0;
Gleam@denominator -> N div Gleam@denominator
end,
Remainder = case erlang:element(3, Loader) of
0 -> 0;
Gleam@denominator@1 -> N rem Gleam@denominator@1
end,
case (Remainder =:= 0) orelse erlang:element(5, Loader) of
true ->
Full;
false ->
Full + 1
end
end.