Composable Callables & Full‑Stack GrowthModel — Practical Guide#

This notebook shows how to:

  1. Wrap scientific models as ``Callable``s (e.g., Elfving 2010 DBH increment, Söderberg 1986 heights, Fridman–Ståhl 2006 mortality, and Edgren–Nylinder 1949 taper + a PriceList).

  2. Build a full‑stack ``GrowthModel`` that runs all components inside grow().

  3. Use the factory to adapt Angle‑Count stands (Bitterlich) to safe working inventories.

  4. Run the DSL controller (triggers/schedules).

  5. Broadcast aggregate runs with the ContextEnsemble (Python/Numba/JAX engines).

All types and utilities come from pyforestry.base.simulation & pyforestry.base.helpers. Everything lives in base/ as per the project conventions.

Notebook Objectives#

  • Explain composable callable patterns and how they connect to GrowthModel runtime orchestration.

  • Provide runnable, copy-safe snippets that work in the docs build environment.

Prerequisites#

  • Python environment with pyforestry installed from this repository.

  • Execute cells in order; random components should use fixed seeds where shown.

Sources#

  • Simulation architecture in pyforestry.base.simulation and pyforestry.simulation.

1) Imports & quick recap#

  • GrowthModel is the factory + behavior.

  • SimulationContext is the sandbox (all mutation + history live here).

  • run_pipeline is the controller: an ordered list of steps over the whole stand, one period at a time.

  • A policy (Callable[[ctx], Sequence[Action]]) is all management is – triggers, schedules and rulesets are policies with an if.

  • AdapterRegistry provides Angle‑Count → pseudo tree‑list/spatial/diameter‑class adapters.

  • ContextEnsemble broadcasts aggregate steps through a batch engine (Python/Numba/JAX).

[1]:
from dataclasses import dataclass
from typing import Any, Callable, Dict, Optional

from pyforestry.base.helpers import (
    PICEA_ABIES,
    PINUS_SYLVESTRIS,
    AngleCount,
    CircularPlot,
    Stand,
    Tree,
)
from pyforestry.base.simulation import (
    ContextEnsemble,
    GrowthModel,
    Requirements,
    Action,
    GrowthStep,
    ManagementStep,
    SimulationContext,
    run_pipeline,
    when,
)

2) Define the scientific components as Callables#

We keep each study/model as a function. In production, you’ll call your real implementations here.

[2]:
# Signatures (type hints are for clarity; not strictly required)
ElfvingDbhIncrementFn = Callable[[Any, float, Optional[float], Any, float, Dict[str, Any]], float]
SoderbergHeightFn     = Callable[[Any, float, Any, Optional[float]], float]
FridmanStahlSurvivalFn= Callable[[Any, float, Optional[float], Any, float, Dict[str, Any]], float]
EdgrenNylinderVolFn   = Callable[[Any, float, float, Any], float]
PriceFn               = Callable[[Any, float, Any], float]

# --- Replace these with your calibrated implementations ---
def elfving_dbh_increment(sp, dbh_cm, h_m, site, dt, state):
    # Δdbh in cm over dt years (placeholder)
    return 0.25 * dt

def soderberg_height(sp, dbh_cm, site, age):
    # height in m from DBH (placeholder)
    return max(1.3, 1.3 + 0.6 * (dbh_cm ** 0.5))

def fridman_stahl_survival(sp, dbh_cm, h_m, site, dt, state):
    # survival fraction in [0,1] for the step (placeholder)
    return max(0.0, min(1.0, 1.0 - 0.006 * dt))

def edgren_nylinder_volume(sp, dbh_cm, h_m, site):
    # taper-based whole-stem volume in m3 (placeholder)
    return 0.00007854 * (dbh_cm ** 2) * h_m

def price_list(sp, vol_m3, site):
    # SEK per m3 (placeholder)
    return 500.0

3) A full‑stack GrowthModel that chains the callables inside grow()#

This model prefers tree_list/spatial inventories. If the input Stand uses Angle‑Count, we’ll ask the factory to adapt to a pseudo tree‑list so per-tree operations work safely.

[3]:
@dataclass
class FullStackCallableModel(GrowthModel):
    dbh_increment_fn: ElfvingDbhIncrementFn
    height_fn:       SoderbergHeightFn
    survival_fn:     FridmanStahlSurvivalFn
    volume_fn:       EdgrenNylinderVolFn
    price_fn:        PriceFn
    remove_zero_weight: bool = True

    def requirements(self) -> Requirements:
        # We need per-tree inventories to run taper/price properly.
        return Requirements(inventory="tree_list")

    def update_step(self, ctx: SimulationContext, dt: float) -> None:
        self.grow(ctx, dt)

    def grow(self, ctx: SimulationContext, dt: float) -> None:
        if ctx.mode not in ("tree_list", "spatial"):
            raise RuntimeError(f"{self.__class__.__name__} requires per-tree mode; got {ctx.mode}.")
        ctx.state["years_since_thin"] = ctx.state.get("years_since_thin", 0.0) + dt
        step_value = 0.0
        for p in ctx.plots:
            for t in p.trees:
                sp = getattr(t, "species", None)
                if sp is None:
                    continue
                dbh = float(getattr(t, "diameter_cm", 0.0) or 0.0)
                h   = float(getattr(t, "height_m",   0.0) or 0.0)
                age = getattr(t, "age", None)

                # 1) DBH increment (Elfving 2010)
                ddbh = float(self.dbh_increment_fn(sp, dbh, (h if h > 0 else None), ctx.site, dt, ctx.state))
                dbh = max(0.0, dbh + ddbh)
                t.diameter_cm = dbh

                # 2) Height update (Söderberg 1986)
                h = float(self.height_fn(sp, dbh, ctx.site, age))
                t.height_m = h

                # 3) Mortality (Fridman–Ståhl 2006)
                surv = float(self.survival_fn(sp, dbh, h, ctx.site, dt, ctx.state))
                surv = 0.0 if surv < 0.0 else (1.0 if surv > 1.0 else surv)
                t.weight_n = float(getattr(t, "weight_n", 1.0)) * surv

                # 4) Value (Edgren–Nylinder 1949 taper + PriceList)
                if dbh > 0.0 and h > 0.0 and t.weight_n > 0.0:
                    vol_m3 = float(self.volume_fn(sp, dbh, h, ctx.site))
                    price  = float(self.price_fn(sp, vol_m3, ctx.site))
                    step_value += vol_m3 * price * float(t.weight_n)

        if self.remove_zero_weight:
            for p in ctx.plots:
                p.trees = [t for t in p.trees if float(getattr(t, "weight_n", 0.0) or 0.0) > 1e-9]

        ctx.state["last_step_value_SEK_per_ha"] = step_value
        ctx.state["cum_value_SEK_per_ha"] = ctx.state.get("cum_value_SEK_per_ha", 0.0) + step_value

4) Build a Stand and run the model through a pipeline#

We’ll use a small synthetic tree‑list Stand. The GrowthModel factory builds a sandboxed SimulationContext for us and run_pipeline steps it.

[4]:
# --- A tiny tree-list stand ---
stand = Stand(
    area_ha=1.0,
    plots=[
        CircularPlot(id=1, area_m2=200.0, trees=[
            Tree(species=PICEA_ABIES, diameter_cm=20.0, height_m=15.0, weight_n=6),
            Tree(species=PICEA_ABIES, diameter_cm=18.0, height_m=13.0, weight_n=5),
        ])
    ],
)

model = FullStackCallableModel(
    dbh_increment_fn=elfving_dbh_increment,
    height_fn=soderberg_height,
    survival_fn=fridman_stahl_survival,
    volume_fn=edgren_nylinder_volume,
    price_fn=price_list,
)

ok, missing = model.can_build(stand, allow_adapters=True, mode_hint="tree_list")
assert ok, f"missing: {missing}"
ctx = model.build_context(stand, mode_hint="tree_list")

run_pipeline(ctx, (GrowthStep(),), years=10.0, step=1.0)
print("Cumulative value (SEK/ha):", ctx.state.get("cum_value_SEK_per_ha", 0.0))
ctx.to_pandas().tail()
Cumulative value (SEK/ha): 7058.081815293103
[4]:
t op details ba_total n_total qmd_total_cm model_state
5 6.0 update_step {'dt': 1.0} 17.706657 530.494635 20.614977 {'t': 6.0, 'years_since_thin': 6.0, 'last_dt':...
6 7.0 update_step {'dt': 1.0} 18.029392 527.311667 20.864689 {'t': 7.0, 'years_since_thin': 7.0, 'last_dt':...
7 8.0 update_step {'dt': 1.0} 18.352762 524.147797 21.114407 {'t': 8.0, 'years_since_thin': 8.0, 'last_dt':...
8 9.0 update_step {'dt': 1.0} 18.676717 521.002910 21.364132 {'t': 9.0, 'years_since_thin': 9.0, 'last_dt':...
9 10.0 update_step {'dt': 1.0} 19.001208 517.876893 21.613863 {'t': 10.0, 'years_since_thin': 10.0, 'last_dt...

5) Angle‑Count stands → pseudo tree‑lists (safe adapter)#

If the Stand carries Bitterlich tallies, you can still run per‑tree logic by opting in to a pseudo tree‑list at build time (your original Stand stays immutable).

[5]:
# --- Bitterlich tally stand ---
ac_plot = CircularPlot(id="ac1", area_m2=10000.0, AngleCount=[
    AngleCount(
        ba_factor=10.0,
        species=[PICEA_ABIES, PINUS_SYLVESTRIS],
        value=[3.0, 2.0],
        # Tallies expand into a pseudo tree list only when the tallied trees'
        # diameters were recorded: stems/ha is derived from them.
        diameters_cm=[[24.0, 21.0, 19.0], [27.0, 23.0]],
    )
])
ac_stand = Stand(area_ha=1.0, plots=[ac_plot])

ok, missing = model.can_build(ac_stand, allow_adapters=True, mode_hint="tree_list")
assert ok, missing
ctx_ac = model.build_context(ac_stand, mode_hint="tree_list")  # auto AC→pseudo tree-list

run_pipeline(ctx_ac, (GrowthStep(),), years=5.0, step=1.0)
print("Inventory origin:", ctx_ac.attrs.get("inventory_origin"))
ctx_ac.to_pandas().head(3)
Inventory origin: angle_count_pseudo_tree_list
[5]:
t op details ba_total n_total qmd_total_cm model_state
0 1.0 update_step {'dt': 1.0} 50.816162 1270.139713 22.569932 {'t': 1.0, 'years_since_thin': 1.0, 'last_dt':...
1 2.0 update_step {'dt': 1.0} 51.633124 1262.518874 22.819195 {'t': 2.0, 'years_since_thin': 2.0, 'last_dt':...
2 3.0 update_step {'dt': 1.0} 52.450774 1254.943761 23.068475 {'t': 3.0, 'years_since_thin': 3.0, 'last_dt':...

6) Using triggers/schedules (DSL)#

You can still add management actions with the DSL. Here we add a simple post trigger that prints a message whenever QMD exceeds a threshold (placeholder for a thinning action).

[6]:
def qmd_exceeds(ctx: SimulationContext, threshold_cm: float = 22.0) -> bool:
    return float(ctx.metrics["QMD"]["TOTAL"]) > threshold_cm

def announce(ctx: SimulationContext):
    print(f"t={ctx.state['t']:.1f} → QMD now {float(ctx.metrics['QMD']['TOTAL']):.2f} cm")

# A trigger is a policy with an `if`: `when(predicate, action)`.
watch = when(
    lambda c: qmd_exceeds(c, 22.0),
    Action(name="qmd_watch", apply=announce),
)

ctx2 = model.build_context(stand, mode_hint="tree_list")
# Placing the ManagementStep after the GrowthStep is what `check_phase="post"`
# used to mean -- the ordering is now the pipeline, so you can read it.
run_pipeline(
    ctx2,
    (GrowthStep(), ManagementStep(watch)),
    years=5.0,
    step=1.0,
)
t=2.0 → QMD now 22.11 cm
t=3.0 → QMD now 22.36 cm
t=4.0 → QMD now 22.61 cm
t=5.0 → QMD now 22.86 cm
[6]:
<pyforestry.base.simulation.core.SimulationContext at 0x7fc10afd40b0>

7) Vectorized aggregate runs with ContextEnsemble#

If you also keep an aggregate approximation of your model, you can opt in to the batch engine by implementing has_batch_engine() and batch_grow_step(ba, n, dt, fert_mask).

[7]:
import numpy as np


class AggregateBANModel(GrowthModel):
    def __init__(self, ba_rel_per_year=0.03, mort_per_year=0.006):
        self.ba_rel = ba_rel_per_year
        self.mort = mort_per_year
    def requirements(self) -> Requirements:
        return Requirements(inventory="aggregate")
    def has_batch_engine(self) -> bool:
        return True
    def batch_grow_step(self, ba: np.ndarray, n: np.ndarray, dt: float, fert_mask: np.ndarray):
        # Simple BA relative growth and mortality on N (vectorized)
        ba2 = ba * (1.0 + self.ba_rel * dt)
        n2  = n  * (1.0 - self.mort * dt)
        return ba2, n2
    def update_step(self, ctx: SimulationContext, dt: float) -> None:
        # Fallback scalar (rarely used when batched)
        ba = float(ctx.metrics["BasalArea"]["TOTAL"]) * (1.0 + self.ba_rel * dt)
        n  = float(ctx.metrics["Stems"]["TOTAL"]) * (1.0 - self.mort * dt)
        ctx.set_aggregate_metrics(ba_total=ba, stems_total=n)

agg_model = AggregateBANModel()

# Build a few contexts in aggregate mode (force with mode_hint)
stands = [stand, stand]  # reuse same stand for brevity
agg_contexts = []
for s in stands:
    ok, _ = agg_model.can_build(s, mode_hint="aggregate")
    ctxa = agg_model.build_context(s, mode_hint="aggregate")
    agg_contexts.append(ctxa)

ens = ContextEnsemble(contexts=agg_contexts, model=agg_model)
for _ in range(5):
    ens.update_step(1.0)

ens.to_pandas().tail()
[7]:
t op details ba_total n_total qmd_total_cm model_state context_id
5 1.0 update_step {'dt': 1.0} 21.248933 499.510751 23.272937 {'t': 1.0, 'years_since_thin': 0.0, 'last_dt':... 1
6 2.0 update_step {'dt': 1.0} 21.886401 496.513686 23.690630 {'t': 2.0, 'years_since_thin': 0.0, 'last_dt':... 1
7 3.0 update_step {'dt': 1.0} 22.542993 493.534604 24.115820 {'t': 3.0, 'years_since_thin': 0.0, 'last_dt':... 1
8 4.0 update_step {'dt': 1.0} 23.219282 490.573397 24.548641 {'t': 4.0, 'years_since_thin': 0.0, 'last_dt':... 1
9 5.0 update_step {'dt': 1.0} 23.915861 487.629956 24.989230 {'t': 5.0, 'years_since_thin': 0.0, 'last_dt':... 1

8) Wrap‑up#

  • Keep each scientific component testable as a Callable.

  • The full‑stack model chains them in one place (grow).

  • The factory adapts inventories (Angle‑Count → pseudo tree‑list) when you ask for per‑tree modes.

  • The DSL orchestrates when things happen; history is logged on every step.

  • For many aggregate scenarios, opt‑in to the batch engine by exposing batch_grow_step.