Skip to content

RFC: pluggable guidance function registry for sampler customization #1124

Description

@FlexOr2

RFC: pluggable guidance function registry for sampler customization

TL;DR

  • Problem. apg_guidance.py already holds 5 guidance variants (APG, CFG, and 3 ADG flavors), but selection is hardcoded in every model and every new APG knob needs its own PR — cf. feat(api): expose sampler_mode and 4 DiT params on /release_task HTTP API #1092 plus a follow-up I've prepared for eta + momentum.
  • Proposal. Adopt a GuidanceFn registry — same DI pattern as PyTorch hooks, Diffusers callback_on_step_end, HF LogitsProcessor, ComfyUI nodes. ~200–300 lines, PyTorch-only, zero behavior change by default.
  • Ask. Interest check before I invest in a Draft PR. "No" is a perfectly fine answer — I'll continue with one-knob-per-PR contributions.

Context

acestep/models/common/apg_guidance.py already ships three core guidance algorithms and two ADG wrapper variants — all in one file:

Function Line Role
apg_forward 33–56 Adaptive Projected Guidance
cfg_forward 59–60 Classic classifier-free guidance
adg_forward 107–188 Angle-based Diffusion Guidance
adg_w_norm_forward 191–212 ADG + norm preservation
adg_wo_clip_forward 215–228 ADG without angle clip

Selection between them is hardcoded inside each model's generate_audio(). For example modeling_acestep_v15_base.py:2025-2039 calls apg_forward in one branch and adg_forward in another, both with fixed parameter choices. MomentumBuffer() is instantiated inline at :1943 with the default momentum=-0.75 baked in.

So the architecture is already halfway to pluggable: multiple variants exist and get selected, but selection logic is frozen and no public entry point lets external users add their own variant (e.g. APG-with-cosine-scheduled-eta, APG-OR, or future paper proposals) without forking and editing each of the four APG-using model files.

The HTTP API inherits this rigidity: every new APG knob needs a dedicated PR — cf. #1092 exposing 5 sampler/DiT params, and a follow-up PR I've prepared for eta + momentum because they're not plumbed through to the HTTP layer either. Each such PR is small but the pattern will repeat indefinitely as new guidance variants are published.

Precedent in this repo

Maintainers have previously done consolidation/DRY work in this area: #1000 ("chore: deduplicate model config and guidance files") and #998 ("chore: deduplicate exact-copy files across model variants"). A registry would be the natural next step on that same consolidation trajectory.

Proposed pattern

Adopt the standard dependency-injection convention used across the Python ML ecosystem:

  • PyTorch: register_forward_hook
  • Diffusers: callback_on_step_end accepts user callables
  • Transformers: StoppingCriteria / LogitsProcessor registries
  • ComfyUI: every node is this pattern, just exposed as a graph

A sketch (signature is a draft for discussion — would refine based on your feedback):

# acestep/models/common/guidance_registry.py (new)
from typing import Protocol, Callable
import torch

class GuidanceFn(Protocol):
    def __call__(
        self,
        pred_cond: torch.Tensor,
        pred_uncond: torch.Tensor,
        guidance_scale: float,
        state: dict,              # owned by the sampler loop, cleared each generate_audio() call
        **params,                 # per-variant params (eta, norm_threshold, angle_clip, etc.)
    ) -> torch.Tensor: ...

_REGISTRY: dict[str, GuidanceFn] = {
    "apg_classic": guidance_apg_classic,   # wraps apg_forward
    "cfg":         guidance_cfg,           # wraps cfg_forward
    "adg":         guidance_adg,           # wraps adg_forward
    "adg_w_norm":  guidance_adg_w_norm,
    "adg_wo_clip": guidance_adg_wo_clip,
}

def register_guidance(name: str) -> Callable[[GuidanceFn], GuidanceFn]:
    """Decorator for third-party variants."""
    ...

HTTP surface (arbitrary callables can't serialize, so use registered-name + param dict):

{
  "guidance_variant": "apg_classic",
  "guidance_params": {"eta": 1.05, "norm_threshold": 1.3, "momentum": -0.25}
}

generate_audio() takes one extra optional keyword: guidance_fn: Callable | None = None. Default preserved → current behavior byte-for-byte unchanged.

Scope estimate

  • PyTorch-only initially. MLX _mlx_apg_forward has a different state-dict shape and no eta parameter; unifying is separate work.
  • Registry file: ~80–120 new lines.
  • Dispatch sites in the 4 APG-using model variants: ~5 lines each → ~20 lines modified.
  • Tests for the registry + each wrapper: ~100–150 lines.
  • Total: ~200–300 lines, 6–8 files. Aligned with "Minimize Blast Radius" if kept PyTorch-only and without touching any existing test.

Honest tradeoffs

I owe an accurate picture of the cost, not just the upside:

Benefits for upstream:

  • Less per-variant maintenance. Community-submitted variants land as small, contained PRs adding one registry entry instead of editing four model files each time.
  • Closes the HTTP-vs-ComfyUI knob-parity gap. HTTP clients stop needing /release_task patches for every new guidance recipe.
  • Research-friendly: your own researchers can prototype new guidance variants without touching the sampler core.

Real costs:

  • API commitment. Once GuidanceFn is public, refactoring the internal signature becomes a semver-breaking change.
  • Debugging blast radius. A user's custom guidance has a bug → looks like inference quality issue → your support triage gets harder. Mitigatable with a "user-registered variant" log stamp.
  • Documentation burden. Per-variant docstrings + recommended param ranges + edge-case warnings. Not huge, but real.
  • Testing surface expansion. apg_classic, cfg, adg, and the two ADG wrappers all become public contracts requiring regression coverage.
  • Two subtle corners that need explicit decisions before code lands:
    1. CFG=1 edge case. When guidance_scale == 1.0, most variants are no-ops but momentum state may still initialize and accumulate. Harmless but wasteful; contract should specify.
    2. State-dict lifetime. Currently MomentumBuffer is scoped to one generate_audio() call. A plugin pattern should preserve that scope explicitly — document that the state dict is fresh per call, never reused across generations.

Prior art

Searched open + closed issues/PRs for guidance, registry, plugin, pluggable, hook. Nothing matched this architectural proposal. Closest related work: #703 (added MLX CFG/APG support — adds a variant, doesn't propose pluggability), #832 (ADG batch bug fix), #1000/#998 (DRY consolidation — supportive precedent). If I missed a prior discussion, happy to close this as duplicate and move the conversation there.

The ask

Are you open in principle to this pattern? I'd rather get a read from you before investing in a draft PR.

  • Yes → I open a Draft PR (~200–300 lines, PyTorch-only) against main for design review. We iterate on signatures / naming / docs together.
  • No / not yet → I drop the proposal and continue contributing incremental HTTP-exposure PRs per variant (feat(api): expose sampler_mode and 4 DiT params on /release_task HTTP API #1092 pattern). That's the status quo and I can live with it.
  • Maybe, but... → your feedback becomes the spec. Anything I have wrong above, any constraint I'm not seeing, please push back.

Happy to refine the design based on whatever you say.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions