anytensor 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- anytensor/__init__.py +208 -0
- anytensor/_version.py +24 -0
- anytensor/backends.py +675 -0
- anytensor/core.py +926 -0
- anytensor/jraph/__init__.py +0 -0
- anytensor/jraph/util.py +0 -0
- anytensor/namespace.py +188 -0
- anytensor/segment.py +465 -0
- anytensor/semantics.py +112 -0
- anytensor/torchscript.py +124 -0
- anytensor/typing.py +104 -0
- anytensor-1.0.0.dist-info/METADATA +86 -0
- anytensor-1.0.0.dist-info/RECORD +16 -0
- anytensor-1.0.0.dist-info/WHEEL +4 -0
- anytensor-1.0.0.dist-info/licenses/LICENSE +21 -0
- anytensor-1.0.0.dist-info/licenses/NOTICE +10 -0
anytensor/__init__.py
ADDED
|
@@ -0,0 +1,208 @@
|
|
|
1
|
+
from . import backends
|
|
2
|
+
from .backends import get_backend
|
|
3
|
+
|
|
4
|
+
# Import einops functions, which work the same way as anytensor.
|
|
5
|
+
from einops import einsum, pack, unpack, rearrange, reduce
|
|
6
|
+
|
|
7
|
+
from .segment import (
|
|
8
|
+
segment_sum,
|
|
9
|
+
segment_max,
|
|
10
|
+
segment_min,
|
|
11
|
+
segment_mean,
|
|
12
|
+
segment_count,
|
|
13
|
+
segment_variance,
|
|
14
|
+
segment_normalize,
|
|
15
|
+
segment_softmax,
|
|
16
|
+
segment_min_or_constant,
|
|
17
|
+
segment_max_or_constant,
|
|
18
|
+
partition_softmax,
|
|
19
|
+
enable_torchscript,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
from .core import (
|
|
23
|
+
repeat,
|
|
24
|
+
take,
|
|
25
|
+
exp,
|
|
26
|
+
log,
|
|
27
|
+
sum,
|
|
28
|
+
min,
|
|
29
|
+
max,
|
|
30
|
+
mean,
|
|
31
|
+
prod,
|
|
32
|
+
shape,
|
|
33
|
+
cumsum,
|
|
34
|
+
reshape,
|
|
35
|
+
transpose,
|
|
36
|
+
concatenate,
|
|
37
|
+
stack,
|
|
38
|
+
maximum,
|
|
39
|
+
minimum,
|
|
40
|
+
sqrt,
|
|
41
|
+
rsqrt,
|
|
42
|
+
where,
|
|
43
|
+
clip,
|
|
44
|
+
astype,
|
|
45
|
+
cast,
|
|
46
|
+
zeros_like,
|
|
47
|
+
ones_like,
|
|
48
|
+
full_like,
|
|
49
|
+
zeros,
|
|
50
|
+
ones,
|
|
51
|
+
full,
|
|
52
|
+
arange,
|
|
53
|
+
matmul,
|
|
54
|
+
inf,
|
|
55
|
+
ninf,
|
|
56
|
+
nan,
|
|
57
|
+
pi,
|
|
58
|
+
e,
|
|
59
|
+
newaxis,
|
|
60
|
+
finfo,
|
|
61
|
+
iinfo,
|
|
62
|
+
dtype,
|
|
63
|
+
is_nan,
|
|
64
|
+
is_finite,
|
|
65
|
+
is_inf,
|
|
66
|
+
isnan,
|
|
67
|
+
isfinite,
|
|
68
|
+
isinf,
|
|
69
|
+
fill_nan,
|
|
70
|
+
nan_fill,
|
|
71
|
+
fill_nan_mask,
|
|
72
|
+
nan_fill_mask,
|
|
73
|
+
nan_to_num,
|
|
74
|
+
equal_nan,
|
|
75
|
+
promote,
|
|
76
|
+
promote_scalars,
|
|
77
|
+
promote_options,
|
|
78
|
+
align_arrays,
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
from .semantics import empty_segment_identity
|
|
82
|
+
from .typing import (
|
|
83
|
+
ArrayT,
|
|
84
|
+
Axes,
|
|
85
|
+
Bool,
|
|
86
|
+
DtypeLike,
|
|
87
|
+
Float,
|
|
88
|
+
FloatArray,
|
|
89
|
+
Inexact,
|
|
90
|
+
Int,
|
|
91
|
+
IntArray,
|
|
92
|
+
Integer,
|
|
93
|
+
Num,
|
|
94
|
+
Real,
|
|
95
|
+
SegmentIds,
|
|
96
|
+
SegmentOut,
|
|
97
|
+
SegmentValues,
|
|
98
|
+
ShapeLike,
|
|
99
|
+
ShapeSize,
|
|
100
|
+
Shaped,
|
|
101
|
+
ShapedArray,
|
|
102
|
+
enable_typecheck,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
try:
|
|
106
|
+
from ._version import __version__
|
|
107
|
+
except ImportError: # pragma: no cover
|
|
108
|
+
__version__ = "0.0.0"
|
|
109
|
+
|
|
110
|
+
__all__ = [
|
|
111
|
+
"backends",
|
|
112
|
+
"get_backend",
|
|
113
|
+
"ArrayT",
|
|
114
|
+
"Axes",
|
|
115
|
+
"Bool",
|
|
116
|
+
"DtypeLike",
|
|
117
|
+
"Float",
|
|
118
|
+
"FloatArray",
|
|
119
|
+
"Inexact",
|
|
120
|
+
"Int",
|
|
121
|
+
"IntArray",
|
|
122
|
+
"Integer",
|
|
123
|
+
"Num",
|
|
124
|
+
"Real",
|
|
125
|
+
"SegmentIds",
|
|
126
|
+
"SegmentOut",
|
|
127
|
+
"SegmentValues",
|
|
128
|
+
"ShapeLike",
|
|
129
|
+
"ShapeSize",
|
|
130
|
+
"Shaped",
|
|
131
|
+
"ShapedArray",
|
|
132
|
+
"enable_typecheck",
|
|
133
|
+
"einsum",
|
|
134
|
+
"pack",
|
|
135
|
+
"unpack",
|
|
136
|
+
"rearrange",
|
|
137
|
+
"reduce",
|
|
138
|
+
"segment_sum",
|
|
139
|
+
"segment_max",
|
|
140
|
+
"segment_min",
|
|
141
|
+
"segment_mean",
|
|
142
|
+
"segment_count",
|
|
143
|
+
"segment_variance",
|
|
144
|
+
"segment_normalize",
|
|
145
|
+
"segment_softmax",
|
|
146
|
+
"segment_min_or_constant",
|
|
147
|
+
"segment_max_or_constant",
|
|
148
|
+
"partition_softmax",
|
|
149
|
+
"enable_torchscript",
|
|
150
|
+
"repeat",
|
|
151
|
+
"take",
|
|
152
|
+
"exp",
|
|
153
|
+
"log",
|
|
154
|
+
"sum",
|
|
155
|
+
"min",
|
|
156
|
+
"max",
|
|
157
|
+
"mean",
|
|
158
|
+
"prod",
|
|
159
|
+
"shape",
|
|
160
|
+
"cumsum",
|
|
161
|
+
"reshape",
|
|
162
|
+
"transpose",
|
|
163
|
+
"concatenate",
|
|
164
|
+
"stack",
|
|
165
|
+
"maximum",
|
|
166
|
+
"minimum",
|
|
167
|
+
"sqrt",
|
|
168
|
+
"rsqrt",
|
|
169
|
+
"where",
|
|
170
|
+
"clip",
|
|
171
|
+
"astype",
|
|
172
|
+
"cast",
|
|
173
|
+
"zeros_like",
|
|
174
|
+
"ones_like",
|
|
175
|
+
"full_like",
|
|
176
|
+
"zeros",
|
|
177
|
+
"ones",
|
|
178
|
+
"full",
|
|
179
|
+
"arange",
|
|
180
|
+
"matmul",
|
|
181
|
+
"inf",
|
|
182
|
+
"ninf",
|
|
183
|
+
"nan",
|
|
184
|
+
"pi",
|
|
185
|
+
"e",
|
|
186
|
+
"newaxis",
|
|
187
|
+
"finfo",
|
|
188
|
+
"iinfo",
|
|
189
|
+
"dtype",
|
|
190
|
+
"is_nan",
|
|
191
|
+
"is_finite",
|
|
192
|
+
"is_inf",
|
|
193
|
+
"isnan",
|
|
194
|
+
"isfinite",
|
|
195
|
+
"isinf",
|
|
196
|
+
"fill_nan",
|
|
197
|
+
"nan_fill",
|
|
198
|
+
"fill_nan_mask",
|
|
199
|
+
"nan_fill_mask",
|
|
200
|
+
"nan_to_num",
|
|
201
|
+
"equal_nan",
|
|
202
|
+
"promote",
|
|
203
|
+
"promote_scalars",
|
|
204
|
+
"promote_options",
|
|
205
|
+
"align_arrays",
|
|
206
|
+
"empty_segment_identity",
|
|
207
|
+
"__version__",
|
|
208
|
+
]
|
anytensor/_version.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
# file generated by vcs-versioning
|
|
2
|
+
# don't change, don't track in version control
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
__all__ = [
|
|
6
|
+
"__version__",
|
|
7
|
+
"__version_tuple__",
|
|
8
|
+
"version",
|
|
9
|
+
"version_tuple",
|
|
10
|
+
"__commit_id__",
|
|
11
|
+
"commit_id",
|
|
12
|
+
]
|
|
13
|
+
|
|
14
|
+
version: str
|
|
15
|
+
__version__: str
|
|
16
|
+
__version_tuple__: tuple[int | str, ...]
|
|
17
|
+
version_tuple: tuple[int | str, ...]
|
|
18
|
+
commit_id: str | None
|
|
19
|
+
__commit_id__: str | None
|
|
20
|
+
|
|
21
|
+
__version__ = version = '1.0.0'
|
|
22
|
+
__version_tuple__ = version_tuple = (1, 0, 0)
|
|
23
|
+
|
|
24
|
+
__commit_id__ = commit_id = None
|