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@core@shape.erl
-module(viva_tensor@core@shape).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/core/shape.gleam").
-export([reshape/2, flatten/1, squeeze/1, squeeze_axis/2, unsqueeze/2, expand_dims/2, take_first/2, take_last/2, slice/3, concat/1, concat_axis/2, stack/2]).
-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).
-file("src/viva_tensor/core/shape.gleam", 22).
?DOC(false).
-spec reshape(viva_tensor@core@tensor:tensor(), list(integer())) -> {ok,
viva_tensor@core@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
reshape(T, New_shape) ->
Old_size = viva_tensor@core@tensor:size(T),
New_size = gleam@list:fold(New_shape, 1, fun(Acc, Dim) -> Acc * Dim end),
case Old_size =:= New_size of
true ->
gleam@result:'try'(
viva_tensor@core@tensor:try_to_list(T),
fun(Data) -> viva_tensor@core@tensor:new(Data, New_shape) end
);
false ->
{error,
{invalid_shape,
<<<<<<<<"Cannot reshape: size mismatch ("/utf8,
(erlang:integer_to_binary(Old_size))/binary>>/binary,
" vs "/utf8>>/binary,
(erlang:integer_to_binary(New_size))/binary>>/binary,
")"/utf8>>}}
end.
-file("src/viva_tensor/core/shape.gleam", 43).
?DOC(false).
-spec flatten(viva_tensor@core@tensor:tensor()) -> viva_tensor@core@tensor:tensor().
flatten(T) ->
Data = viva_tensor@core@tensor:to_list(T),
viva_tensor@core@tensor:from_list(Data).
-file("src/viva_tensor/core/shape.gleam", 49).
?DOC(false).
-spec squeeze(viva_tensor@core@tensor:tensor()) -> viva_tensor@core@tensor:tensor().
squeeze(T) ->
Data = viva_tensor@core@tensor:to_list(T),
New_shape = gleam@list:filter(
viva_tensor@core@tensor:shape(T),
fun(D) -> D /= 1 end
),
Final_shape = case New_shape of
[] ->
[1];
_ ->
New_shape
end,
case viva_tensor@core@tensor:new(Data, Final_shape) of
{ok, Result} ->
Result;
{error, _} ->
T
end.
-file("src/viva_tensor/core/shape.gleam", 432).
?DOC(false).
-spec list_at(list(AJWH), integer()) -> {ok, AJWH} | {error, nil}.
list_at(Lst, Index) ->
case Index < 0 of
true ->
{error, nil};
false ->
_pipe = Lst,
_pipe@1 = gleam@list:drop(_pipe, Index),
gleam@list:first(_pipe@1)
end.
-file("src/viva_tensor/core/shape.gleam", 63).
?DOC(false).
-spec squeeze_axis(viva_tensor@core@tensor:tensor(), integer()) -> {ok,
viva_tensor@core@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
squeeze_axis(T, Axis) ->
Shp = viva_tensor@core@tensor:shape(T),
case list_at(Shp, Axis) of
{error, _} ->
{error, {dimension_error, <<"Axis out of bounds"/utf8>>}};
{ok, D} ->
case D =:= 1 of
false ->
{error,
{invalid_shape, <<"Dimension at axis is not 1"/utf8>>}};
true ->
Data = viva_tensor@core@tensor:to_list(T),
New_shape = begin
_pipe = Shp,
_pipe@1 = gleam@list:index_map(
_pipe,
fun(Dim, I) -> {Dim, I} end
),
_pipe@2 = gleam@list:filter(
_pipe@1,
fun(Pair) -> erlang:element(2, Pair) /= Axis end
),
gleam@list:map(
_pipe@2,
fun(Pair@1) -> erlang:element(1, Pair@1) end
)
end,
viva_tensor@core@tensor:new(Data, New_shape)
end
end.
-file("src/viva_tensor/core/shape.gleam", 85).
?DOC(false).
-spec unsqueeze(viva_tensor@core@tensor:tensor(), integer()) -> viva_tensor@core@tensor:tensor().
unsqueeze(T, Axis) ->
Data = viva_tensor@core@tensor:to_list(T),
Shp = viva_tensor@core@tensor:shape(T),
Rnk = erlang:length(Shp),
Insert_at = case Axis < 0 of
true ->
(Rnk + Axis) + 1;
false ->
Axis
end,
{Before, After} = gleam@list:split(Shp, Insert_at),
New_shape = lists:append([Before, [1], After]),
case viva_tensor@core@tensor:new(Data, New_shape) of
{ok, Result} ->
Result;
{error, _} ->
T
end.
-file("src/viva_tensor/core/shape.gleam", 103).
?DOC(false).
-spec expand_dims(viva_tensor@core@tensor:tensor(), integer()) -> viva_tensor@core@tensor:tensor().
expand_dims(T, Axis) ->
unsqueeze(T, Axis).
-file("src/viva_tensor/core/shape.gleam", 110).
?DOC(false).
-spec take_first(viva_tensor@core@tensor:tensor(), integer()) -> viva_tensor@core@tensor:tensor().
take_first(T, N) ->
Data = viva_tensor@core@tensor:to_list(T),
case viva_tensor@core@tensor:shape(T) of
[] ->
T;
[First_dim | Rest_dims] ->
Take_n = gleam@int:min(N, First_dim),
Stride = gleam@list:fold(Rest_dims, 1, fun(Acc, D) -> Acc * D end),
New_data = gleam@list:take(Data, Take_n * Stride),
New_shape = [Take_n | Rest_dims],
case viva_tensor@core@tensor:new(New_data, New_shape) of
{ok, Result} ->
Result;
{error, _} ->
T
end
end.
-file("src/viva_tensor/core/shape.gleam", 128).
?DOC(false).
-spec take_last(viva_tensor@core@tensor:tensor(), integer()) -> viva_tensor@core@tensor:tensor().
take_last(T, N) ->
Data = viva_tensor@core@tensor:to_list(T),
case viva_tensor@core@tensor:shape(T) of
[] ->
T;
[First_dim | Rest_dims] ->
Take_n = gleam@int:min(N, First_dim),
Stride = gleam@list:fold(Rest_dims, 1, fun(Acc, D) -> Acc * D end),
Skip = (First_dim - Take_n) * Stride,
New_data = gleam@list:drop(Data, Skip),
New_shape = [Take_n | Rest_dims],
case viva_tensor@core@tensor:new(New_data, New_shape) of
{ok, Result} ->
Result;
{error, _} ->
T
end
end.
-file("src/viva_tensor/core/shape.gleam", 442).
?DOC(false).
-spec list_at_float(list(float()), integer()) -> {ok, float()} | {error, nil}.
list_at_float(Lst, Index) ->
list_at(Lst, Index).
-file("src/viva_tensor/core/shape.gleam", 422).
?DOC(false).
-spec compute_strides(list(integer())) -> list(integer()).
compute_strides(Shape) ->
Reversed = lists:reverse(Shape),
{Strides, _} = gleam@list:fold(
Reversed,
{[], 1},
fun(Acc, Dim) ->
{S, Running} = Acc,
{[Running | S], Running * Dim}
end
),
Strides.
-file("src/viva_tensor/core/shape.gleam", 413).
?DOC(false).
-spec multi_to_flat(list(integer()), list(integer())) -> integer().
multi_to_flat(Indices, Shape) ->
Strides = compute_strides(Shape),
_pipe = gleam@list:zip(Indices, Strides),
gleam@list:fold(
_pipe,
0,
fun(Acc, Pair) ->
{Idx, Stride} = Pair,
Acc + (Idx * Stride)
end
).
-file("src/viva_tensor/core/shape.gleam", 401).
?DOC(false).
-spec flat_to_multi(integer(), list(integer())) -> list(integer()).
flat_to_multi(Flat, Shape) ->
Reversed = lists:reverse(Shape),
{Indices, _} = gleam@list:fold(
Reversed,
{[], Flat},
fun(Acc, Dim) ->
{Idxs, Remaining} = Acc,
Idx = case Dim of
0 -> 0;
Gleam@denominator -> Remaining rem Gleam@denominator
end,
Next = case Dim of
0 -> 0;
Gleam@denominator@1 -> Remaining div Gleam@denominator@1
end,
{[Idx | Idxs], Next}
end
),
Indices.
-file("src/viva_tensor/core/shape.gleam", 484).
?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/core/shape.gleam", 480).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/core/shape.gleam", 148).
?DOC(false).
-spec slice(viva_tensor@core@tensor:tensor(), list(integer()), list(integer())) -> {ok,
viva_tensor@core@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
slice(T, Start, Lengths) ->
gleam@result:'try'(
viva_tensor@core@tensor:try_to_list(T),
fun(Data) ->
Shp = viva_tensor@core@tensor:shape(T),
R = viva_tensor@core@tensor:rank(T),
case (erlang:length(Start) =:= R) andalso (erlang:length(Lengths)
=:= R) of
false ->
{error,
{dimension_error,
<<"Slice dimensions must match tensor rank"/utf8>>}};
true ->
case R of
1 ->
gleam@result:'try'(
begin
_pipe = list_at(Start, 0),
gleam@result:map_error(
_pipe,
fun(_) ->
{dimension_error,
<<"Slice start is missing"/utf8>>}
end
)
end,
fun(S) ->
gleam@result:'try'(
begin
_pipe@1 = list_at(Lengths, 0),
gleam@result:map_error(
_pipe@1,
fun(_) ->
{dimension_error,
<<"Slice length is missing"/utf8>>}
end
)
end,
fun(Len) ->
Sliced = begin
_pipe@2 = Data,
_pipe@3 = gleam@list:drop(
_pipe@2,
S
),
gleam@list:take(_pipe@3, Len)
end,
viva_tensor@core@tensor:new(
Sliced,
[Len]
)
end
)
end
);
_ ->
New_size = gleam@list:fold(
Lengths,
1,
fun(Acc, D) -> Acc * D end
),
Result_data = begin
_pipe@4 = range_int(0, New_size - 1),
gleam@list:fold(
_pipe@4,
{ok, []},
fun(Acc@1, Flat_idx) ->
gleam@result:'try'(
Acc@1,
fun(Values) ->
Local_indices = flat_to_multi(
Flat_idx,
Lengths
),
Global_indices = gleam@list:map2(
Local_indices,
Start,
fun(L, S@1) -> L + S@1 end
),
Global_flat = multi_to_flat(
Global_indices,
Shp
),
gleam@result:'try'(
begin
_pipe@5 = list_at_float(
Data,
Global_flat
),
gleam@result:map_error(
_pipe@5,
fun(_) ->
{index_out_of_bounds,
Global_flat,
erlang:length(
Data
)}
end
)
end,
fun(Value) ->
{ok, [Value | Values]}
end
)
end
)
end
)
end,
gleam@result:'try'(
Result_data,
fun(Result_data@1) ->
Result = lists:reverse(Result_data@1),
viva_tensor@core@tensor:new(Result, Lengths)
end
)
end
end
end
).
-file("src/viva_tensor/core/shape.gleam", 211).
?DOC(false).
-spec concat(list(viva_tensor@core@tensor:tensor())) -> viva_tensor@core@tensor:tensor().
concat(Tensors) ->
Data = gleam@list:flat_map(
Tensors,
fun(T) -> viva_tensor@core@tensor:to_list(T) end
),
viva_tensor@core@tensor:from_list(Data).
-file("src/viva_tensor/core/shape.gleam", 454).
?DOC(false).
-spec concat_axis_value(
list(viva_tensor@core@tensor:tensor()),
integer(),
list(integer())
) -> {ok, float()} | {error, viva_tensor@core@error:tensor_error()}.
concat_axis_value(Tensors, Tensor_idx, Local_indices) ->
gleam@result:'try'(
begin
_pipe = list_at(Tensors, Tensor_idx),
gleam@result:map_error(
_pipe,
fun(_) ->
{dimension_error,
<<"Concat index does not map to a source tensor"/utf8>>}
end
)
end,
fun(T) ->
gleam@result:'try'(
viva_tensor@core@tensor:try_to_list(T),
fun(Data) ->
Strides = compute_strides(viva_tensor@core@tensor:shape(T)),
Local_flat = begin
_pipe@1 = gleam@list:zip(Local_indices, Strides),
gleam@list:fold(
_pipe@1,
0,
fun(Acc, Pair) ->
{Idx, Stride} = Pair,
Acc + (Idx * Stride)
end
)
end,
_pipe@2 = list_at_float(Data, Local_flat),
gleam@result:map_error(
_pipe@2,
fun(_) ->
{index_out_of_bounds,
Local_flat,
erlang:length(Data)}
end
)
end
)
end
).
-file("src/viva_tensor/core/shape.gleam", 446).
?DOC(false).
-spec concat_data(list(viva_tensor@core@tensor:tensor())) -> {ok, list(float())} |
{error, viva_tensor@core@error:tensor_error()}.
concat_data(Tensors) ->
gleam@list:fold(
Tensors,
{ok, []},
fun(Acc, T) ->
gleam@result:'try'(
Acc,
fun(Values) ->
gleam@result:'try'(
viva_tensor@core@tensor:try_to_list(T),
fun(Data) -> {ok, lists:append(Values, Data)} end
)
end
)
end
).
-file("src/viva_tensor/core/shape.gleam", 222).
?DOC(false).
-spec concat_axis(list(viva_tensor@core@tensor:tensor()), integer()) -> {ok,
viva_tensor@core@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
concat_axis(Tensors, Axis) ->
case Tensors of
[] ->
{error, {invalid_shape, <<"Cannot concatenate empty list"/utf8>>}};
[Single] ->
{ok, Single};
[First | Rest] ->
Base_shape = viva_tensor@core@tensor:shape(First),
R = erlang:length(Base_shape),
case (Axis >= 0) andalso (Axis < R) of
false ->
{error,
{dimension_error,
<<"Invalid axis for concatenation"/utf8>>}};
true ->
Shapes_ok = gleam@list:all(
Rest,
fun(T) ->
T_shape = viva_tensor@core@tensor:shape(T),
case erlang:length(T_shape) =:= R of
false ->
false;
true ->
_pipe = gleam@list:zip(Base_shape, T_shape),
_pipe@1 = gleam@list:index_map(
_pipe,
fun(Pair, I) -> {Pair, I} end
),
gleam@list:all(
_pipe@1,
fun(X) ->
{{Dim_a, Dim_b}, I@1} = X,
(I@1 =:= Axis) orelse (Dim_a =:= Dim_b)
end
)
end
end
),
case Shapes_ok of
false ->
{error,
{invalid_shape,
<<"Shapes must match except on concat axis"/utf8>>}};
true ->
Concat_dim = gleam@list:fold(
Tensors,
0,
fun(Acc, T@1) ->
case begin
_pipe@2 = gleam@list:drop(
viva_tensor@core@tensor:shape(T@1),
Axis
),
gleam@list:first(_pipe@2)
end of
{ok, D} ->
Acc + D;
{error, _} ->
Acc
end
end
),
New_shape = begin
_pipe@3 = Base_shape,
gleam@list:index_map(
_pipe@3,
fun(D@1, I@2) -> case I@2 =:= Axis of
true ->
Concat_dim;
false ->
D@1
end end
)
end,
case Axis =:= 0 of
true ->
gleam@result:'try'(
concat_data(Tensors),
fun(Data) ->
viva_tensor@core@tensor:new(
Data,
New_shape
)
end
);
false ->
Total_size = gleam@list:fold(
New_shape,
1,
fun(Acc@1, D@2) -> Acc@1 * D@2 end
),
Result_data = begin
_pipe@4 = range_int(0, Total_size - 1),
gleam@list:fold(
_pipe@4,
{ok, []},
fun(Acc@2, Flat_idx) ->
gleam@result:'try'(
Acc@2,
fun(Values) ->
Indices = flat_to_multi(
Flat_idx,
New_shape
),
Axis_idx = case begin
_pipe@5 = gleam@list:drop(
Indices,
Axis
),
gleam@list:first(
_pipe@5
)
end of
{ok, I@3} ->
I@3;
{error, _} ->
0
end,
{Tensor_idx,
Local_axis_idx,
_} = gleam@list:fold(
Tensors,
{-1, Axis_idx, 0},
fun(Acc@3, T@2) ->
{Found_t,
Remaining,
T_idx} = Acc@3,
case Found_t >= 0 of
true ->
Acc@3;
false ->
T_axis_size = case begin
_pipe@6 = gleam@list:drop(
viva_tensor@core@tensor:shape(
T@2
),
Axis
),
gleam@list:first(
_pipe@6
)
end of
{ok,
D@3} ->
D@3;
{error,
_} ->
0
end,
case Remaining
< T_axis_size of
true ->
{T_idx,
Remaining,
T_idx};
false ->
{-1,
Remaining
- T_axis_size,
T_idx
+ 1}
end
end
end
),
Local_indices = begin
_pipe@7 = Indices,
gleam@list:index_map(
_pipe@7,
fun(Idx, I@4) ->
case I@4 =:= Axis of
true ->
Local_axis_idx;
false ->
Idx
end
end
)
end,
gleam@result:'try'(
concat_axis_value(
Tensors,
Tensor_idx,
Local_indices
),
fun(Value) ->
{ok,
[Value |
Values]}
end
)
end
)
end
)
end,
gleam@result:'try'(
Result_data,
fun(Result_data@1) ->
Result = lists:reverse(
Result_data@1
),
viva_tensor@core@tensor:new(
Result,
New_shape
)
end
)
end
end
end
end.
-file("src/viva_tensor/core/shape.gleam", 357).
?DOC(false).
-spec stack(list(viva_tensor@core@tensor:tensor()), integer()) -> {ok,
viva_tensor@core@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
stack(Tensors, Axis) ->
case Tensors of
[] ->
{error, {invalid_shape, <<"Cannot stack empty list"/utf8>>}};
[First | Rest] ->
Base_shape = viva_tensor@core@tensor:shape(First),
Shapes_ok = gleam@list:all(
Rest,
fun(T) -> viva_tensor@core@tensor:shape(T) =:= Base_shape end
),
case Shapes_ok of
false ->
{error, {shape_mismatch, Base_shape, []}};
true ->
_ = erlang:length(Tensors),
R = erlang:length(Base_shape),
Insert_axis = case Axis < 0 of
true ->
(R + Axis) + 1;
false ->
Axis
end,
case (Insert_axis >= 0) andalso (Insert_axis =< R) of
false ->
{error,
{dimension_error,
<<"Invalid axis for stacking"/utf8>>}};
true ->
Unsqueezed = gleam@list:map(
Tensors,
fun(T@1) -> unsqueeze(T@1, Insert_axis) end
),
concat_axis(Unsqueezed, Insert_axis)
end
end
end.