Writing a vocabulary II: extend it by inheritance

Writing a vocabulary II: extend it by inheritance#

Add to a vocabulary the ordinary Python way: subclass it.

Time: ~5 min · Runs on: CPU (same code on GPU) · You need: a family of kernels

class Perturbed(Orbit, slug="perturbed") carries Orbit’s whole vocabulary forward and adds its own entries; a kernel built on Perturbed sees both. Re-declaring an inherited entry is fine if it matches exactly, and refuses – by name – if it does not.

Hide code cell source

import sys, pathlib
sys.path.insert(0, str(pathlib.Path.cwd().parent / "_shared"))
from nb_helpers import plot_style
plot_style()
from hawk import Param, Vector
from hawk.ext import DATA_ONLY, KernelKind


class Orbit(KernelKind, slug="orbit"):
    r: Vector[3]
    v: Vector[3]
    dt: Param
    g: Param
    guard = DATA_ONLY


Orbit
__main__.Orbit

Extend it by subclassing:

from hawk import Scalar


class Perturbed(Orbit, slug="perturbed"):
    """Adds ONE entry, `cd` -- a per-sample drag coefficient."""
    cd: Scalar


print("Perturbed's vocabulary:", [name for name, _ in Perturbed.kind.vocabulary])
Perturbed's vocabulary: ['r', 'v', 'dt', 'g', 'cd']

drag_fall, built on Perturbed, sees r, v, dt, g (Orbit’s own) AND cd (its own) without re-declaring any of them:

import numpy as np
from hawk import Mutable
from hawk.math import norm, vec


@Perturbed
def drag_fall(r, v, dt, g, cd, r_next: Mutable[Vector[3]], v_next: Mutable[Vector[3]]):
    a = vec(0.0, 0.0, -g) - cd * norm(v) * v
    v_new = v + dt * a
    v_next = v_new
    r_next = r + dt * v_new


drag_fall
<Kernel drag_fall slots=7>

hawk.artifact.build + hawk.load/hawk.run build and run it on the host:

import pathlib, tempfile

import hawk.artifact

work_dir = pathlib.Path(tempfile.mkdtemp())
n = 3
r0 = np.zeros((3, n)); r0[2] = 2.0
v0 = np.zeros((3, n)); v0[0] = 3.0
r_next, v_next = np.zeros((3, n)), np.zeros((3, n))
hawk.artifact.build(drag_fall, work_dir, targets=("host",))
hawk.run(hawk.load(work_dir, "drag_fall"), r=r0, v=v0, dt=0.1, g=9.81,
       cd=np.array([0.0, 0.01, 0.05]), r_next=r_next, v_next=v_next)
print("x after one step, by drag coefficient 0 / 0.01 / 0.05:", r_next[0])
x after one step, by drag coefficient 0 / 0.01 / 0.05: [0.3    0.2991 0.2955]

When a re-declaration disagrees#

Equal re-declarations are harmless (a subclass restating a base’s own entry verbatim); a DIFFERENT one for the same name refuses, naming the entry, the base and both declarations – never a silent shadow:

from hawk.ir import HawkError

try:
    class BadPerturbed(Orbit, slug="bad_perturbed"):
        r: Scalar  # Orbit already declares `r` as a Vector[3] -- disagrees
except HawkError as exc:
    print("refused:", exc)
refused: class 'BadPerturbed'(KernelKind, ...): vocabulary entry 'r' is declared differently by 'Orbit' (Plane(form='in', ttype=TensorType(shape=(3,), dtype='f64', tag=None), op=None, dims=(), staged=False)) and 'BadPerturbed' (Plane(form='in', ttype=TensorType(shape=(), dtype='f64', tag=None), op=None, dims=(), staged=False)) — an inherited entry may only be re-declared with an EQUAL resolved declaration

The instance-form twin#

Kind.extend does the same merge without a class statement – the same object the class form built, for the same inputs:

extended = Orbit.kind.extend("perturbed", vocabulary={"cd": Scalar})
extended == Perturbed.kind
True

What just happened#

  • Perturbed(Orbit, slug=...) merged Orbit’s vocabulary with its own ONE new entry (cd); drag_fall used both without re-declaring anything.

  • A conflicting re-declaration refuses by NAME, at class-creation time – before any kernel built from it can ever run.

  • Kind.extend(...) is the same merge, spelled on an instance instead of a class statement – one mechanism, two spellings.

Try this#

Add a second child, class Lofted(Orbit, slug="lofted"): lift: Scalar, and a kernel that uses lift alongside r, v, dt, g.

Next#

Your own function, your own derivative – give a kernel on Perturbed a custom primitive. Deeper: the inheritance rules in hawk.ext.KernelKind’s own docstring (API reference).