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