Vectorization with vmap and lift¶
Write a calculation for one value, then apply it to many values with vmap.
Use lift when a function combines two values and you want a reduction,
running accumulation, outer product, or broadcasting call. Both reuse the
scalar function's arithmetic and error handling.
Choose the interface for the task¶
| Task | Start with |
|---|---|
Apply a built-in function, such as exp, to a batch |
batch.exp(values) |
| Apply your own function independently to elements, rows, or columns | vmap[f](...) |
| Combine values with a binary function or evaluate all pairs | lift[f](...) |
| Compute a standard total, product, or dot product | batch.sum, batch.prod, or batch.dot |
Here vectorization means applying a scalar or slice calculation across arrays. Eligible work may use CPU threads; arbitrary-precision numbers do not become native SIMD lanes. The scalar numerical guarantees still apply.
Start with one scalar calculation¶
This example computes squared distances from a center. The same function handles one position, a vector of positions, or every position-center pair. There is no separate batch implementation of the formula.
"""Map a scalar calculation, then combine or pair its results."""
from apn_mojo import Batch, Integer, batch, integer, lift, vmap
def squared_distance(value: Integer, center: Integer) raises -> Integer:
var difference = value - center
return difference * difference
def row_total(row: Batch[Integer]) raises -> Integer:
return batch.sum(row)
def main() raises:
var positions = Batch[Integer]([1, 2, 4])
var distances = vmap[squared_distance](in_axes=(0, None))(positions, Integer(2))
print("squared distances from 2:", distances)
var add = lift[integer.add](identity=Integer(0), associative=True)
print("total:", add.reduce(distances, axis=None))
print("running totals:", add.accumulate(distances))
var centers = Batch[Integer]([0, 2])
print("all pairs:", lift[squared_distance]().outer(positions, centers))
var grid = Batch[Integer]([1, 2, 3, 4, 5, 6]).reshape([2, 3])
print("row totals:", vmap[row_total](in_axes=0)(grid))
print("column totals:", vmap[row_total](in_axes=1)(grid))
Run from the repository root pixi run mojo run -I src docs/examples/vectorization.mojo
Output
squared distances from 2: [1, 0, 4]
total: 5
running totals: [1, 1, 5]
all pairs: [[1, 1], [4, 0], [16, 4]]
row totals: [6, 15]
column totals: [5, 7, 9]
In vmap[squared_distance](in_axes=(0, None)), brackets select the function
at compile time. The first call configures its mapping; the next supplies
the data. Axis 0 visits the positions, while None shares the center
across every call. A shared argument can also be a whole batch.
Pass changing settings as explicit arguments rather than capturing local
variables in a closure. For native variables, construct the intended APN
type, such as Integer(center). Arbitrary-width integer literals can be
passed directly and remain exact.
For library functions, select a family declaration such as integer.add
or float.sqrt. Root functions such as apn_mojo.add are overloaded
dispatchers and cannot serve as one compile-time function value. A mapped
context= is forwarded to the scalar function when its signature accepts it.
Map rows or columns¶
Mapping an axis removes that axis from each input seen by the function.
For a [2, 3] matrix, in_axes=0 passes two rows of shape [3], and
in_axes=1 passes three columns of shape [2]. That is why row_total
accepts a Batch[Integer], even though each call returns one Integer.
A function taking a scalar Integer cannot consume a whole row. Nest mappings
to reach individual entries of a higher-rank batch; the
composition example shows how.
For built-in elementwise work across any rank, batch.* already handles it.
out_axes chooses where the new mapped axis goes in the output. A Bool
result becomes a Mask; a tuple produces a corresponding tuple of outputs.
Batch results must have a consistent shape. Empty mappings with batch
results need out_shape when no call can establish that shape. See
rows, columns, and tensor results
for a runnable example of these options.
Combine values with lift¶
The lifted addition in the example supports these calls:
| Operation | Meaning |
|---|---|
reduce(values, axis=None) |
Combine every element into one scalar |
reduce(values, axis=0) |
Combine along one axis; this is the default axis |
accumulate(values, axis=0) |
Return every running prefix along an axis |
outer(a, b) |
Evaluate every pair, with shape a.shape() + b.shape() |
| A direct call on two arguments | Apply the binary function over broadcast-compatible shapes |
Use keepdims=True to retain reduced dimensions with size one. An empty
reduction needs an identity configured on the lifted function or an
initial value supplied to the reduction. For nonempty input, initial
is combined before the selected elements.
associative=True promises that regrouping the function's calls leaves
the result unchanged. Exact integer addition satisfies this promise, so
long reductions can combine partial totals in parallel. It does not require
commutativity: operand order is preserved. Leave the default False for
subtraction and other order-sensitive functions. A fold combines values of
one result family; outer products and broadcasting also accept supported
functions with different input and result types.
Keep the shape rules separate¶
vmap pairs entries along mapped axes, whose extents must match. It does
not repeat a one-element vector to match a longer mapped axis. Pass a scalar
to share one value instead. batch.* elementwise functions and lifted
binary calls align trailing dimensions and allow size-one dimensions.
Batch operators have an additional restriction on two rank-one vectors:
their lengths must match.
"""The vector rules for operators, mapped functions and lift."""
from apn_mojo import ArithmeticContext, Batch, FloatFormat, Integer, batch, integer, lift, vmap
def main() raises:
var values = Batch[Integer]([1, 2, 3])
var singleton = Batch[Integer]([10])
try:
_ = values + singleton
except:
print("operator: vector lengths must match")
print("batch.add:", batch.add(values, singleton))
try:
_ = vmap[integer.add]()(values, singleton)
except:
print("vmap: mapped extents must match")
print("lift:", lift[integer.add]()(values, singleton))
var context = ArithmeticContext(format=FloatFormat(24))
print("batch.add with context:", batch.add(values, singleton, context=context))
print("batch.atan2 shape:", batch.atan2(values, singleton, context=context).shape())
Run from the repository root pixi run mojo run -I src docs/examples/batch_broadcasting.mojo
Output
operator: vector lengths must match
batch.add: [11, 12, 13]
vmap: mapped extents must match
lift: [11, 12, 13]
batch.add with context: [11.0, 12.0, 13.0]
batch.atan2 shape: [3]
Masks used for selection must match the selected shape. batch.where
broadcasts its two value arguments to the mask's shape; the mask itself
does not broadcast.
Choose how a fold rounds¶
Integer and Rational folds remain exact. Float and Complex reductions and
accumulations keep intermediate values exact by default, then round each
returned total or prefix once. context= controls that final rounding.
This is useful for sums and products, but an intermediate such as 1/3
cannot be held exactly in a finite binary representation and raises.
Family functions can also take mixed exact inputs. In this example, the
Rational inputs 1/3 and 2/3 reach the first addition exactly, so their
sum is 1 before it is rounded to a Float. Ball accumulation instead carries
an enclosure from one step to the next.
"""Exact numeric inputs and Ball folds use the scalar function's rules."""
from apn_mojo import Batch, Integer, Rational, FloatFormat, ArithmeticContext, BallContext, lift
from apn_mojo import rational, float, ball
def main() raises:
var integers = Batch[Integer]([1, 2, 3])
print(lift[rational.add]()(integers, Rational(1, 3)))
print(lift[rational.add]().outer(10, integers))
var thirds = Batch[Rational]([Rational(1, 3), Rational(2, 3)])
var format = ArithmeticContext(format=FloatFormat(24))
# The operands stay exact; only the returned total rounds.
print(lift[float.add]().reduce(thirds, axis=None, context=format))
var add = lift[ball.add]()
var prefixes = add.accumulate(thirds, context=BallContext(precision=24))
print(ball.contains(prefixes[0], Rational(1, 3)))
print(prefixes[1].is_exact(), ball.contains(prefixes[1], 1))
Run from the repository root pixi run mojo run -I src docs/examples/lift_inputs.mojo
Output
[4/3, 7/3, 10/3]
[11, 12, 13]
1.0
True
True True
Use lift[f](exact=False) when rounding after each step is the calculation
you intend. A rounded addition is generally not associative, so do not
promise associativity for that mode. Ball folds carry an enclosure through
each call and receive BallContext; exact does not change their behavior.
The lifting reference includes a runnable comparison of exact and stepwise rounding, as well as empty reductions and in-place accumulators.
Write functions that can run independently¶
Mapped functions must be pure: avoid printing, shared mutation, random state, and other effects whose outcome depends on call order. A failed element may be called again to obtain its error message. The library discards partial output and reports the lowest failing logical index.
Eligible long mappings can use the worker pool; short runs and some signatures stay on the calling thread. Set the thread limit as shown in thread controls. Parallel execution does not relax numerical guarantees or the associativity promise. The mapping reference lists signature limits and supported result types.