Compute

Marrow provides SIMD-vectorized compute kernels for arithmetic, comparisons, selection and casting. All kernels are null-aware; the exact null handling depends on the operation.

String, temporal, boolean and conditional kernels are implemented in Mojo but not yet bound to Python — they are reachable from a query plan instead.

The compute functions live in marrow.compute, and their names and signatures follow pyarrow.compute for the functions it implements — mc.add, mc.subtract, mc.cast, mc.filter. It is not a drop-in replacement: some options (skip_nulls=False, nan_is_null=True, multi-key sort_keys) raise NotImplementedError, and memory_pool is ignored.

Arithmetic

Element-wise binary operations: add, subtract, multiply, divide.

a = ma.array([1, 2, 3, 4, 5], type=ma.int64())
b = ma.array([10, 20, 30, 40, 50], type=ma.int64())

print("add:     ", mc.add(a, b))
print("subtract:", mc.subtract(b, a))
print("multiply:", mc.multiply(a, b))
print("divide:  ", mc.divide(b, a))
add:      PrimitiveArray[int64]([11, 22, 33, 44, 55])
subtract: PrimitiveArray[int64]([9, 18, 27, 36, 45])
multiply: PrimitiveArray[int64]([10, 40, 90, 160, 250])
divide:   PrimitiveArray[int64]([10, 10, 10, 10, 10])

Floats work the same way:

x = ma.array([1.0, 2.0, 3.0])
y = ma.array([0.5, 1.5, 2.5])
print("add:   ", mc.add(x, y))
print("divide:", mc.divide(x, y))
add:    PrimitiveArray[float64]([1.5, 3.5, 5.5])
divide: PrimitiveArray[float64]([2.0, 1.3333333333333333, 1.2])
Notemarrow.compute compares like pyarrow; marrow.expr compares like SQL

The two answer differently on NaN, deliberately. marrow.compute mirrors pyarrow.compute, so its comparisons are IEEE — a NaN equals nothing, itself included. marrow.expr is a SQL engine, and DuckDB, DataFusion and Polars all agree that SQL’s comparisons are total:

expression pc.equal / pc.greater col("a") == … / > …
nan = nan False True
nan <> nan True False
nan > 1.0 False True

So a NaN is its own GROUP BY group, its own join key and its own rank peer, and it sorts after inf. NULL is unaffected and stays three-valued in both.

Float division by zero answers a value — 10.0 / 0.0 is inf, 0.0 / 0.0 is nan — while integer // and % by zero answer NULL, as SQL does. The two genuinely want different answers: division by zero has a value in the reals’ completion and integer division by zero does not.

Null propagation

If either operand at a position is null, the result at that position is null. This mirrors SQL’s three-valued logic.

a = ma.array([1, None, 3, None])
b = ma.array([10, 20, 30, 40])

result = mc.add(a, b)
print(result)    # index 1 and 3 are null
print("null count:", result.null_count)
PrimitiveArray[int64]([11, NULL, 33, NULL])
null count: 2
# Both inputs can contribute nulls
a = ma.array([None, 2, 3, None])
b = ma.array([10, None, 30, None])

print(mc.add(a, b))   # null at 0, 1, and 3
PrimitiveArray[int64]([NULL, NULL, 33, NULL])

Aggregates

Most aggregates are not eager kernel calls (boolean any and all, below, are the exception). sum, mean, min, max, product, count, count_distinct and the rest reduce a column inside a query plan, so that a whole-table reduction and a GROUP BY are the same code path. Wrap a batch with ma.memtable(), describe the reduction, and .collect() runs it:

batch = ma.record_batch({"v": ma.array([1, None, 3, None, 5])})
q = ma.memtable(batch)

out = q.aggregate(
    total=("sum", "v"),
    smallest=("min", "v"),
    largest=("max", "v"),
    average=("mean", "v"),
)
print(out.collect().to_pylist())
[{'total': 9, 'smallest': 1, 'largest': 5, 'average': 3.0}]

Nulls are skipped, so sum above is 1 + 3 + 5. Adding a by= key groups rather than reducing to one row — the only change is the argument:

sales = ma.record_batch({
    "region": ma.array(["east", "west", "east", "west"]),
    "amount": ma.array([10, 20, 30, 40]),
})
by_region = ma.memtable(sales).aggregate(
    by=["region"], total=("sum", "amount"),
).order_by("region")
print(by_region.collect().to_pylist())
[{'region': 'east', 'total': 40}, {'region': 'west', 'total': 60}]

Accumulator types widen

sum and product do not preserve the input type — integers accumulate in int64 and floats in float64, so a long column cannot silently overflow its own width. mean is always float64:

ints = ma.record_batch({"v": ma.array([10, 20, 30], type=ma.int32())})
res  = ma.memtable(ints).aggregate(total=("sum", "v"), avg=("mean", "v")).collect()
print("sum of int32 ->", res.column("total").type)
print("mean         ->", res.column("avg").type)
sum of int32 -> int64
mean         -> float64

Counting

count counts non-null values; ma.count_star() counts rows, which differs on a nullable column. count_distinct is exact, while approx_count_distinct uses a fixed-size HyperLogLog sketch — far cheaper in memory on a high-cardinality column:

ids = ma.record_batch({"id": ma.array([1, 1, 2, 3, 3, 3, None, 4])})
print(ma.memtable(ids).aggregate(
    rows=ma.count_star(),
    values=("count", "id"),
    distinct=("count_distinct", "id"),
    approx=("approx_count_distinct", "id"),
).collect().to_pylist())
[{'rows': 8, 'values': 7, 'distinct': 4, 'approx': 4}]

Boolean aggregates

any and all operate on boolean arrays. Nulls are skipped.

flags = ma.array([True, False, None, True])
print("any:", mc.any(flags))   # True  — at least one True
print("all:", mc.all(flags))   # False — False is present
any: True
all: False
all_true = ma.array([True, True, None])
print("all with nulls:", mc.all(all_true))   # True — only True values

all_false = ma.array([None, None], type=ma.bool_())
print("all empty/null:", mc.all(all_false))   # True (identity)
print("any empty/null:", mc.any(all_false))   # False (identity)
all with nulls: True
all empty/null: True
any empty/null: False

Selection

filter

filter(array, mask) keeps elements where the boolean mask is True. The mask must be a BoolArray of the same length.

arr  = ma.array([10, 20, 30, 40, 50])
mask = ma.array([True, False, True, False, True])

print(mc.filter(arr, mask))   # [10, 30, 50]
PrimitiveArray[int64]([10, 30, 50])

Nulls in the source array are preserved through the filter:

arr  = ma.array([1, None, 3, 4, None])
mask = ma.array([True, True, False, True, True])

print(mc.filter(arr, mask))   # [1, NULL, 4, NULL]
PrimitiveArray[int64]([1, NULL, 4, NULL])

Nulls in the mask are treated as False (the element is excluded):

arr  = ma.array([1, 2, 3, 4])
mask = ma.array([True, None, True, None])

print(mc.filter(arr, mask))   # [1, 3]
PrimitiveArray[int64]([1, 3])

drop_null

drop_null(array) removes all null positions. Equivalent to filtering by the array’s own validity bitmap.

arr = ma.array([1, None, 3, None, 5])
print(mc.drop_null(arr))   # [1, 3, 5]
PrimitiveArray[int64]([1, 3, 5])

Works on numeric array types:

f = ma.array([1.0, None, 3.0, None, 5.0])
print(mc.drop_null(f))
PrimitiveArray[float64]([1.0, 3.0, 5.0])

Comparisons

Element-wise comparison kernels return a boolean array. Nulls propagate: if either input at a position is null, the output is null.

a = ma.array([1, 2, 3, None, 5])
b = ma.array([1, 3, 2, 4,    5])

print("equal:        ", mc.equal(a, b))
print("not_equal:    ", mc.not_equal(a, b))
print("less:         ", mc.less(a, b))
print("less_equal:   ", mc.less_equal(a, b))
print("greater:      ", mc.greater(a, b))
print("greater_equal:", mc.greater_equal(a, b))
equal:         BoolArray([True, False, False, NULL, True])
not_equal:     BoolArray([False, True, True, NULL, False])
less:          BoolArray([False, True, False, NULL, False])
less_equal:    BoolArray([True, True, False, NULL, True])
greater:       BoolArray([False, False, True, NULL, False])
greater_equal: BoolArray([True, False, True, NULL, True])

The null at index 3 propagates to all output arrays:

result = mc.equal(a, b)
print("null count:", result.null_count)      # 1
print("is_valid(3):", result[3].is_valid())     # False
null count: 1
is_valid(3): False

Casting

cast converts an array to another type. By default a lossy conversion raises; pass safe=False for the raw truncating/wrapping conversion.

ints = ma.array([1, 2, 3], type=ma.int32())
print("to float64:", mc.cast(ints, ma.float64()))

# safe=False truncates toward zero instead of raising
floats = ma.array([1.9, -1.9, 2.5], type=ma.float64())
print("truncated: ", mc.cast(floats, ma.int32(), safe=False))
to float64: PrimitiveArray[float64]([1.0, 2.0, 3.0])
truncated:  PrimitiveArray[int32]([1, -1, 2])

Casts cover the numeric, boolean, string/binary and temporal families, including decimal rescaling and dictionary decode at the Mojo level.

Null behaviour summary

Operation Null input Null output
add, subtract, multiply, divide either operand null null propagates
equal, less, greater, … either operand null null propagates
sum, product, min, max, mean element null element skipped
count_distinct, approx_count_distinct element null element skipped
any, all element null element skipped
filter (source null) source element null null preserved in output
filter (mask null) mask element null element excluded
drop_null element null element removed
cast element null null preserved
Back to top