hawk.math

hawk.math#

hawk.math — the free-function namespace a kernel body writes math through.

A namespace, not a layer: every name here either is one of hawk.trace’s free functions, re-exported unchanged, or a two-line composition of them, so hawk.math and hawk.ir.ops declare the same vocabulary by construction (tests/test_math_namespace.py traces and emits every name in __all__).

It exists because exp, log, tanh, vec and max want spelling without importing forty names one at a time, and because three spellings have no free function elsewhere in HAWK yet:

  • sample_index() — the lane’s own index as a value;

  • split_index() — the (major, minor) decomposition of a flattened two-dimensional launch domain;

  • take() — the free-function spelling of a plane’s at() read.

The names match the previous code generator’s vmath where it has one, so a re-authored body moves across with its diffs readable; semantics do not move (hawk.math.max is HAWK’s maximum).

Importing this module imports the tracer, so import hawk never reaches it: it is asked for by name, from hawk import math as m.

Element-wise math, on host and device alike, with numpy’s names:

  • exp / log: exp exp2 expm1 log log2 log10 log1p power (pow, **) sqrt rsqrt cbrt hypot

  • trigonometric: sin cos tan asin acos atan atan2

  • hyperbolic: sinh cosh tanh asinh acosh atanh

  • rounding: floor ceil trunc round rint

  • misc: absolute (abs) sign copysign fmod remainder fdim fma clip minimum (min) maximum (max) isnan isinf isfinite

  • special: erf erfc

Semantics differing from numpy, by design: round is C’s (halves away from zero; rint is numpy’s half-to-even np.round), and minimum/maximum ignore a NaN operand (np.fmin/np.fmax). remainder is numpy’s floor-mod (sign of the divisor), fmod C’s (sign of the dividend). The class tests return bool.

Shapes: every element-wise function takes vector and matrix operands as the arithmetic operators do — equal shapes, or one operand rank-0 broadcast — and refuses any other combination in the operators’ words. exp log sqrt rsqrt, the trigonometric functions but atan2, tanh, abs, power, minimum and maximum lower to one aether expression at any rank; the rest, with the comparisons and land/lor/lnot, are mapped entry by entry at trace time and re-assembled — aether spells them at rank 0 only — so a class test or a comparison on a matrix is a bool matrix, the mask where()/select() takes per entry (a rank-0 branch broadcasts).

Derivatives: the rounding family, sign and the class tests carry a zero derivative. abs takes +1 at zero (as JAX); copysign is |a| times b’s sign, constant in b; fmod/remainder differentiate as a - trunc(a/b)*b / a - floor(a/b)*b; fdim is zero at a == b; hypot is zero at the origin; clip passes the gradient to x on the closed interval [lo, hi] (torch’s clamp; JAX halves it at a tie), to hi above it or whenever lo > hi, and to lo below it.

Not yet: lgamma/tgamma/digamma, Bessel functions, and integer-specific operations beyond the existing ones.

Module Attributes

Functions

abs(…)

|x| (np.absolute; also abs(x)).

absolute(…)

np.absolute/np.power: the abs(x)/x ** y operators as functions.

acos(…)

Elementwise acos of a traced value.

acosh(…)

Inverse hyperbolic cosine, x >= 1 (np.arccosh).

argmax(…)

The index (rank-0 i32) of the first maximum: of a rank-1 value (argmax(logits)), of a Python list of rank-0 values, or of rank-0 components given positionally (argmax(a, b, c)).

as_pure(…)

Elementwise as_pure of a traced value.

as_vec3(…)

Elementwise as_vec3 of a traced value.

asin(…)

Elementwise asin of a traced value.

asinh(…)

Inverse hyperbolic sine (np.arcsinh).

atan(…)

Elementwise atan of a traced value.

atan2(…)

The kinds aether spells at rank 0 only, mapped entry by entry on a vector or matrix like the functions below.

atanh(…)

Inverse hyperbolic tangent, |x| <= 1 (np.arctanh).

cbrt(…)

Real cube root, odd in x (np.cbrt).

ceil(…)

Smallest integer >= x (np.ceil).

clip(…)

minimum(maximum(x, lo), hi) with NaN propagated from any argument (np.clip).

copysign(…)

|a| with the sign bit of b (np.copysign).

cos(…)

Elementwise cos of a traced value.

cosh(…)

Hyperbolic cosine (np.cosh).

cross(…)

cross of two traced values.

dispatch(…)

branches[clamp(kind, 0)] — the finite per-element dispatch.

dot(…)

dot of two traced values.

erf(…)

The error function (scipy.special.erf).

erfc(…)

1 - erf(x), accurate for large x (scipy.special.erfc).

exp(…)

Elementwise exp of a traced value.

exp2(…)

2**x (np.exp2).

expm1(…)

exp(x) - 1, accurate near 0 (np.expm1).

fdim(…)

max(a - b, 0) (C fdim).

floor(…)

Piecewise-constant (lever 2): the largest integer <= x, as an f64 value.

fma(…)

a * b + c with one rounding (C fma).

fmod(…)

C remainder of a / b, sign of a (np.fmod).

hypot(…)

sqrt(a*a + b*b) without overflow (np.hypot).

isfinite(…)

True where x is neither inf nor NaN (np.isfinite); bool-typed, component-wise on a vector or matrix.

isinf(…)

True where x is +-inf (np.isinf); bool-typed, component-wise on a vector or matrix.

isnan(…)

True where x is NaN (np.isnan); bool-typed, component-wise on a vector or matrix.

land(…)

Logical and of two bool values, component-wise.

lnot(…)

Logical not of a bool value, component-wise.

log(…)

Elementwise log of a traced value.

log10(…)

Base-10 logarithm (np.log10).

log1p(…)

log(1 + x), accurate near 0 (np.log1p).

log2(…)

Base-2 logarithm (np.log2).

lor(…)

Logical or of two bool values, component-wise.

max(…)

max of two traced values.

maximum(…)

max of two traced values.

min(…)

min of two traced values.

minimum(…)

min of two traced values.

n_samples()

The readable sample count: the TRUE nSamples, never count.

norm(…)

Elementwise norm of a traced value.

outer(…)

The outer product u v^T, a FREE function because Python has no infix for it and @ is already the contraction (mv/mm): the contracting dot and expanding outer must be spelled apart.

pow(…)

a ** b (np.power; also the ** operator).

power(…)

a ** b (np.power; also the ** operator).

quat_conj(…)

Elementwise quat_conj of a traced value.

quat_mul(…)

quat_mul of two traced values.

quat_recip(…)

Elementwise quat_recip of a traced value.

quat_rotate(…)

quat_rotate of two traced values.

random_bernoulli(…)

A Bernoulli draw as a real 1.0 / 0.0: u < p (aether's rule).

random_exponential(…)

An exponential draw with rate (mean 1/rate): -log(1 - u) / rate, not aether's own -log(u) / rate, since u in [0, 1) keeps the argument in (0, 1] so no lane draws +inf.

random_lognormal(…)

exp(mu + sigma * z) — aether's Generator::lognormal(mu, sigma).

random_multivariate_normal(…)

A correlated normal vector mean + L z with L the (lower) Cholesky factor of the covariance and z D independent standard draws on sub-counters counter*D + k — aether's LowerTriangular composition, through the mv op.

random_normal(…)

A normal draw N(mean, sd^2) — at the defaults exactly the primitive op (hawk.trace.random_normal()); otherwise mean + sd * z.

random_poisson1(…)

A Poisson(1)-distributed count per lane, as a real value: Knuth's product method, multiplying uniform draws until the running product falls below e**-1 — the number of products still above it is the draw.

random_uniform(…)

A uniform draw on [lo, hi) — at the defaults exactly the primitive op, otherwise the composition lo + (hi - lo) * u, aether's own formula.

random_uniform_int(…)

An integer draw uniform over the closed range [lo, hi], carried as a real value: floor(lo + (hi - lo + 1) * u) — u < 1 strictly, so hi is reached and never exceeded (aether's uniformInt convention).

remainder(…)

Floor-mod a - floor(a / b) * b, sign of b; a zero result carries b's sign (np.remainder, not C's IEEE remainder).

rint(…)

Nearest integer, halves to even (np.rint).

round(…)

Nearest integer, halves AWAY from zero (C round; np.round rounds halves to even, which is rint()).

rsqrt(…)

Elementwise rsqrt of a traced value.

sample_index()

This lane's GLOBAL sample index, as a value (the own(i)): it is base + flat under a partition, so an expression built on it is partition-invariant — safe for the oracle to compare — and lets a 2-D-domain kernel recover the coordinates a flattened launch folded together.

select(…)

cond ? a: b — what the ast pass rewrites control flow into.

sign(…)

-1, 0 or 1 by the sign of x; NaN stays NaN and both zeros give +0 (np.sign).

sin(…)

Elementwise sin of a traced value.

sinh(…)

Hyperbolic sine (np.sinh).

split_index(…)

Split this lane's index into its (major, minor) components.

sqrt(…)

Elementwise sqrt of a traced value.

take(…)

Read buf at an absolute index — the free-function spelling of at.

tan(…)

Elementwise tan of a traced value.

tanh(…)

Elementwise tanh of a traced value.

transpose(…)

Elementwise transpose of a traced value.

trunc(…)

Integer part, toward zero (np.trunc).

vec(…)

Build a rank-1 value from scalar components: both vec(a, b, c) and vec([a, b, c]) work, the second being what a static comprehension evaluates to.

vsum(…)

Elementwise sum of a traced value.

where(…)

Branchless cond ? a: b — the np.where spelling of select().