Metadata-Version: 2.4
Name: grafted
Version: 0.1.0
Summary: Adapter-first composition for turning frozen video models into conditioned models
License: MIT
Requires-Python: >=3.10
Description-Content-Type: text/markdown
Requires-Dist: torch>=2.1
Requires-Dist: PyYAML>=6.0
Requires-Dist: h5py>=3.16.0
Provides-Extra: dev
Requires-Dist: pytest>=8; extra == "dev"
Requires-Dist: ruff>=0.6; extra == "dev"
Provides-Extra: wandb
Requires-Dist: wandb>=0.17; extra == "wandb"

# grafted

Adapter-first composition for turning **frozen** video models — diffusion or flow — into conditioned world models.

```
prediction = compose( base(x_t, t, cond), adapter(x_t, t, cond) )
```

The base model never trains. You implement four things; the training loop is handed to you.

| You implement | You get |
|---|---|
| `BaseVideoModel` — `encode` / `decode` / `denoise` / `generate` / `prepare_batch` | your backbone drives training *and* its own native sampler |
| `Adapter` — the trainable part | composition, conditioning dropout, checkpointing |
| `Composer` — how the adapter folds onto the base | `Add`, `GatedResidual`, `MaskMix`, `Replace`, stackable wrappers |
| `Objective` — what the loss means | `Diffusion`, `FlowMatching`, masked losses |

## Train something

```python
from grafted import AdapterExperiment, GatedResidual, FlowMatching

exp = AdapterExperiment(
    base      = MyBackbone.from_pretrained(path),   # frozen for you
    adapter   = MyAdapter(),
    composer  = GatedResidual(),
    objective = FlowMatching(),
)
exp.fit(dataloader, steps=10_000)
```

or from YAML — `AdapterExperiment.from_yaml("run.yaml").fit(dataloader)`:

```yaml
name: my-run
base:      {type: wan22_ti2v,     config: {checkpoint: /ckpt/wan22}}
adapter:   {type: hyperalign,     config: {rank: 8}}
composer:  {type: gated_residual, config: {gate_bias: 0.0}}
objective: {type: flow_matching,  config: {shift: true}}
training:  {steps: 20000, learning_rate: 1.0e-4, amp_dtype: bf16}
callbacks:
  - {type: console, config: {every_n_steps: 10}}
```

The outer schema is fixed and tiny. Everything under `config:` is validated by the component's **own** dataclass, so a typo raises at build time instead of silently doing nothing — the failure mode that quietly invalidates a run.

## Conditioning is declared, not dug out

An adapter says what it reads. The backbone says what it supplies. They are matched **when the model is constructed**:

```python
class MyAdapter(Adapter):
    consumes = (Condition("action", aliases=("act",), dim=4),)

    def forward(self, x_t, t, cond, base_output=None):
        cond.action     # [B, A] — alias resolved, frames collapsed,
                        # dropout applied, dtype and device matched
```

```python
class MyBackbone(BaseVideoModel):
    def provides(self):
        return ("act", "fs", "c_concat", "c_crossattn")
```

A mismatch fails before a GPU is touched, and the message is the fix:

```
MyAdapter consumes condition 'depth' (looked for 'depth'), but
MyBackbone.provides() provides 'act', 'c_concat', 'c_crossattn', 'fs'.
Add the key in the backbone's prepare_batch, declare an alias on the
Condition, or mark it required=False.
```

This exists because the alternative is silent. Conditioning used to travel as an untyped dict, so a key the backbone spelled differently produced a zeros tensor — and an adapter training on nothing has a perfectly healthy loss curve. `per_frame` (`"mean"` / `"sum"` / `"keep"`) is explicit for the same reason: mean-vs-sum changes what the model is conditioned on, so neither is a silent default.

Presence is checked at build time; feature width is checked on first use, because a backbone cannot know the action width without a batch. Backbone-native keys stay reachable by subscript (`cond["fs"]`). To call an adapter outside training — a probe, an ablation — use `adapter.view({...})` to get the same prepared inputs.

## Composition is a component, not an enum

Composition rules stack, so orthogonal ideas stay orthogonal:

```python
NoiseGated(PretrainWarmup(GatedResidual(), steps=500), lo=0.3, hi=1.0)
```

* `NoiseGated` — scale the correction by noise level; the adapter only acts where the outcome is still undecided.
* `PretrainWarmup` — train the adapter branch alone first, so the gate has something worth choosing between.

## Callbacks, including the ones components bring themselves

The trainer does forward, backward, accumulate, clip, step, checkpoint. Every measurement is a callback — and a component can declare its own:

```python
class MyAdapter(Adapter):
    def callbacks(self):
        return (MyGateProbe(self), ActionSensitivity(self))
```

The trainer collects those automatically from the backbone, the adapter, the composer and the objective. An adapter's instrumentation travels with the adapter; nobody writing a training script has to know it exists.

## Extending from another package

Register components from *your* distribution — no fork, no edit here:

```toml
[project.entry-points."grafted.adapters"]
hyperalign = "my_research.adapters:HyperAlignAdapter"

[project.entry-points."grafted.backbones"]
wan22_ti2v = "my_research.backbones:Wan22TI2V"
```

`type: hyperalign` then resolves in any grafted config. Groups: `grafted.backbones`, `grafted.adapters`, `grafted.composers`, `grafted.objectives`, `grafted.callbacks`, `grafted.datasets`. Resolution is lazy, so installing plugins costs nothing at import.

Datasets register the same way — a `data:` block in the config is a `type` plus a `config` block like any other component, so a dataset's own knobs (an HDF5 path, a clip window, an action dimension) live next to it instead of scattered across CLI flags:

```yaml
data: {type: metaworld, config: {path: ds/metaworld_corner2.hdf5, window: 16, stride: 4}}
```

```python
from grafted.registry import register

@register("dataset", "metaworld")
class MetaWorldClips(Dataset): ...
```

This is also how private research code stays private: it depends on `grafted`, `grafted` never depends on it, and a test enforces the direction.

### Worked examples: the bundled plugins

Three complete plugins ship in this repo, each a backbone plus a reference `*_latent_residual` adapter — and each showing a different *way* to wrap a real model without grafted ever depending on it:

| Plugin | Backbone | Wraps | How |
|---|---|---|---|
| `plugins/dynamicrafter` | [DynamiCrafter](https://github.com/Doubiiu/DynamiCrafter) | a vendored copy of upstream's own `lvdm` tree | delegates to it: `prepare_batch` → `get_batch_input`, `denoise` → `apply_model`, `generate` → `DDIMSampler` |
| `plugins/wan22` | Wan2.2 TI2V-5B | diffusers' native `WanPipeline` | no vendored code at all — diffusers has shipped Wan2.2 since 0.35, so this just drives its pipeline components directly |
| `plugins/easyanimate` | EasyAnimate | diffusers' native `EasyAnimatePipeline` | same shape as `wan22`; the one wrinkle is its CFG doubles the batch before the transformer sees it, which `generate` has to split around `compose_fn` |

Vendoring vs. wrapping diffusers are both "delegate, don't reimplement" — which one fits depends only on whether diffusers already carries the model. Neither touches `grafted` itself: `type: dynamicrafter` / `type: wan22_ti2v` / `type: easyanimate` all resolve purely through each plugin's own entry points.

```bash
pip install -e . && pip install -e plugins/dynamicrafter
python plugins/dynamicrafter/examples/train_dc.py --checkpoint /path/to/dynami512.ckpt
```

`wan22` and `easyanimate` have no bespoke smoke-test script — they run through the same generic `examples/train.py` every backbone does, which is itself the point: swap a config, not a script.

```bash
pip install -e plugins/wan22
python examples/train.py --config examples/configs/wan22.yaml --steps 20
```

In all three, the adapter is injected by wrapping the single call every prediction flows through (`apply_model` for DynamiCrafter, `transformer.forward` for the diffusers-native pair) — never a reimplemented sampling loop.

## Setup

Core dependencies are `torch` and `PyYAML`. Backbone- and metric-specific requirements belong to whoever provides those backbones.

### Fresh environment

```bash
cd grafted
python -m venv .venv && source .venv/bin/activate
pip install -e ".[dev]"                      # grafted + pytest + ruff
pytest tests                                 # 30 tests, ~2s, CPU only
```

That is the whole library. All three plugins are optional and heavy:

```bash
pip install -e plugins/dynamicrafter         # lvdm deps: open-clip, transformers, kornia, ...
pytest plugins/dynamicrafter/tests           # 9 pass, 5 skip without a checkpoint

pip install -e plugins/wan22                 # diffusers, transformers, accelerate, sentencepiece
pytest plugins/wan22/tests                   # CPU tests pass, 4 skip without a checkpoint

pip install -e plugins/easyanimate           # diffusers, transformers, accelerate
pytest plugins/easyanimate/tests             # CPU tests pass, 4 skip without a checkpoint
```

### With uv (workspace)

Plugins are separate distributions but live in this repo, so they are declared as uv workspace members — that is what makes `grafted-dynamicrafter`, `grafted-wan22` and `grafted-easyanimate` resolvable from a local path instead of PyPI. Whether they are *installed* is then controlled by the `plugins` dependency group:

```bash
uv sync                      # core only — torch + PyYAML, no plugin
uv sync --group plugins      # + grafted-dynamicrafter, grafted-wan22, grafted-easyanimate
uv run --group plugins pytest tests plugins/dynamicrafter/tests plugins/wan22/tests plugins/easyanimate/tests
```

**A bare `uv sync` uninstalls the plugin again** — that is the group working as intended (core stays lean), not a bug. Pass `--group plugins` consistently when working on plugin code, or add it to `[tool.uv] default-groups` if you always want it.

### Reusing an existing torch environment

Skip dependency resolution entirely — useful when torch/CUDA is already pinned:

```bash
uv pip install -e . --no-deps
uv pip install -e plugins/dynamicrafter --no-deps   # or plugins/wan22, plugins/easyanimate
```

## Testing

Run from the repo root; both packages are importable once installed, so no `PYTHONPATH` is needed.

```bash
pytest tests                                 # core:   30 tests, CPU, ~2s
pytest plugins/dynamicrafter/tests           # plugin: 9 CPU tests (+5 gated)
pytest plugins/wan22/tests                   # plugin: 10 CPU tests (+4 gated)
pytest plugins/easyanimate/tests             # plugin: 10 CPU tests (+4 gated)
pytest tests plugins/dynamicrafter/tests plugins/wan22/tests plugins/easyanimate/tests   # all
```

The gated tests in each plugin need real weights and a GPU, and skip cleanly unless you point at a checkpoint — same pattern, one env var per plugin:

```bash
GRAFTED_DC_CHECKPOINT=/path/to/dynami512.ckpt pytest plugins/dynamicrafter/tests -q
GRAFTED_WAN22_CHECKPOINT=/path/to/wan22-diffusers pytest plugins/wan22/tests -q
GRAFTED_EASYANIMATE_CHECKPOINT=/path/to/easyanimate-diffusers pytest plugins/easyanimate/tests -q
```

Each set covers the load (every non-EMA/expected weight must match), a `denoise` step, identity-at-init composition against the real frozen base, and a full training step. Expect ~100 s per plugin — most of it loading the checkpoint. `wan22` and `easyanimate` take a diffusers-format checkpoint directory or Hub repo id, not the raw upstream release layout — see each plugin's `backbone.py` module docstring if `from_pretrained` complains about a missing `model_index.json`.

### End-to-end smoke run

```bash
python plugins/dynamicrafter/examples/train_dc.py \
    --checkpoint /path/to/dynami512.ckpt --steps 10
```

Random clips, so the loss is meaningless — the point is that config → entry-point resolution → checkpoint load → conditioning → composed forward → backward → step all runs against real weights. ~1 s/step on an RTX 3090, 2.9M trainable against a 2609M frozen base.

`wan22` and `easyanimate` have no bespoke script for this — the same smoke test is `python examples/train.py --config examples/configs/wan22.yaml --steps 10` (or `easyanimate.yaml`), which is itself the thing being demonstrated: the generic entrypoint didn't need to change for a diffusers-native backbone.

### What the tests actually check

| File | Covers |
|---|---|
| `tests/test_core.py` | composition (incl. zero-init identity, gate-cap freeze guard, stacked composers), objectives and loss masking, callback cadence/dedup/collection, the training loop, generation, config validation |
| `tests/test_boundary.py` | `grafted` imports no research code or heavy backbone deps; every module imports with only torch + PyYAML; `__all__` matches a frozen snapshot |
| `plugins/dynamicrafter/tests/` | entry-point registration, the dependency arrow, adapter behaviour, dynamic rescale, and the real-weights integration set |
| `plugins/wan22/tests/` | entry-point registration, the dependency arrow, per-token timestep expansion, adapter behaviour, and the real-weights integration set |
| `plugins/easyanimate/tests/` | entry-point registration, the dependency arrow, the CFG-doubled-batch `compose_fn` split, adapter behaviour, and the real-weights integration set |

A failure in `test_public_surface_is_frozen` means the public API changed — update the snapshot deliberately, don't reflexively.

## Troubleshooting

**`KeyError: Unknown backbone 'dynamicrafter'. Registered: dummy.`** (or `'wan22_ti2v'`, or `'easyanimate'`)

The plugin is not installed *in the environment you are running from*. Components resolve through installed distribution metadata, so a plugin that is importable on `PYTHONPATH` but not installed will not be found — deliberately, since that is also what stops a half-installed plugin from half-working.

Check which environment you are actually in, and what it can see:

```bash
python -c "
import importlib.metadata as md
print(sorted(d.metadata['Name'] for d in md.distributions() if (d.metadata['Name'] or '').startswith('grafted')))
print([e.name for e in md.entry_points(group='grafted.backbones')])"
```

An empty entry-point list is the answer. Fix with `uv sync --group plugins`, or `pip install -e plugins/dynamicrafter` in that environment.

This bites most often when a repo has more than one venv — an activated `VIRTUAL_ENV` from another project does not override uv's project environment, and uv says so in a warning that is easy to scroll past.

## Status

Early. The interfaces above are the ones being stabilised; expect churn below them until 0.2.

**Planned, not yet implemented:** LoRA- and ControlNet-style adapters. Both need nothing new from the `Adapter` interface — `attach_base_model` (a hook for reaching into the frozen base's own modules) and `reuses_base_output=False` (declaring that the base's own forward pass now depends on the adapter, so a cached output can't be reused) already exist for exactly this shape of adapter; neither has a real user yet. `src/grafted/adapters/lora.py` and `src/grafted/adapters/controlnet.py` are stubs — registered under `lora`/`controlnet`, construction always raises `NotImplementedError` naming what's missing, rather than silently doing nothing or half-working. See [Writing an adapter](docs/guides/adapter.md).
