x = jnp.arange(12, dtype=jnp.float32)
xDive into Deep Learning · §1.1
Storing & transforming data with tensors
The n-dimensional arrays that every model in this book is built on.
Motivation
ndarray.Rank = number of axes; shape = size per axis.
01
Getting Started
creating & inspecting tensors
Getting Started
arange(n) builds a 1-D tensor of evenly spaced values:
Array([ 0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11.], dtype=float32)
(12,)
numel() → total elements. shape → size along each axis. We ask for float32 because nearly all neural-net math is in floating point.
Getting Started
For weight init, randn draws from \mathcal{N}(0, 1):
Array([[ 1.6226422 , 2.0252647 , -0.43359444, -0.07861735],
[ 0.1760909 , -0.97208923, -0.49529874, 0.4943786 ],
[ 0.6643493 , -0.9501635 , 2.1795304 , -1.9551506 ]], dtype=float32)
Also zeros, ones, full(shape, value), eye(n). Random values break symmetry when initializing network weights; lists let you type a tensor by hand.
Getting Started
02
Indexing & Slicing
reading & writing elements, rows, ranges
Indexing & Slicing
Indexing & Slicing
JAX arrays can’t be mutated..at[i].set(v) returns a new array:
Array([[ 0., 1., 2., 3.],
[ 4., 5., 17., 7.],
[ 8., 9., 10., 11.]], dtype=float32)
Indexing & Slicing
03
Operations
elementwise math, joins, comparisons, broadcasting
Operations
The operators + - * / ** act elementwise on matching shapes:
(Array([ 3., 4., 6., 10.], dtype=float32),
Array([-1., 0., 2., 6.], dtype=float32),
Array([ 2., 4., 8., 16.], dtype=float32),
Array([0.5, 1. , 2. , 4. ], dtype=float32),
Array([ 1., 4., 16., 64.], dtype=float32))
Any scalar→scalar map (exp, sin, log) extends to a whole tensor.
Operations
Operations
Comparisons return a boolean tensor.
A ready-made mask:
Array([[False, True, False, True],
[False, False, False, False],
[False, False, False, False]], dtype=bool)
==, <, > build masks; sum, mean, max collapse axes; add dim= to reduce just one.
Operations · the exception
Size-1 axes are virtually stretched
a 3\times1 plus a 1\times2 gives a 3\times2:
Array([[0, 1],
[1, 2],
[2, 3]], dtype=int32)
Any axis of size 1 stretches to match the other tensor, without a copy.
Compatible only if each axis is equal or 1.
Operations · the exception
Line up (3, 2) and (2, 3) from the right, pairing 2 with 3 and 3 with 2: no pair matches, neither member is 1, so the framework raises rather than guessing:
add got incompatible shapes for broadcasting: (3, 2), (2, 3).
Broadcasting aligns shapes from the right; each axis pair must be equal or 1.
04
Memory & Interop
in-place updates and leaving the tensor world
Y = Y + XPerformance
Every arithmetic expression allocates a new tensor
costly when Y is gigabytes and updated many times per second:
False
id(Y) changed: Y is now bound to a new tensor object.
Performance
JAX has no in-place write; a functional update returns a new array, so id changes:
False
Under jit, XLA fuses such updates and reuses buffers, recovering the in-place benefit.
Interop
Convert to / from a NumPy ndarray:
(numpy.ndarray, jaxlib._jax.ArrayImpl)
The result is a copy; host/device arrays don’t share storage here.
Wrap-up
arange, zeros, ones, randn, tensor([…])..shape, .numel(), reshape.cat.X[:] = …, +=), or in JAX via jit buffer reuse..item() for scalars.