feat(sampling): add contextual diffusion sampler

This commit is contained in:
Artificial Sweetener
2026-08-02 02:43:26 -04:00
parent 24492d97a3
commit 29ee772b1d
31 changed files with 2420 additions and 209 deletions
@@ -10,7 +10,6 @@
from __future__ import annotations
from collections.abc import Sequence
from importlib import import_module
from types import ModuleType
from typing import Any, cast
@@ -18,10 +17,8 @@ from typing import Any, cast
import torch
from ..domain.tiled_diffusion import (
LatentTile,
TiledDiffusionPlan,
build_tiled_diffusion_plan,
gaussian_tile_weights,
)
from ..shared.logging import get_logger
from . import sampling_samplers, sampling_schedulers
@@ -31,11 +28,8 @@ from .tiled_sampling import (
ApplyModel,
Latent,
ModelFunctionWrapper,
SemanticTileWeightCache,
make_tiled_model_args,
new_spatial_weight_buffer,
TilePredictionAccumulator,
reject_unsupported_conditioning,
spatial_tile_slicer,
validate_latent_samples,
validate_sampling_controls,
validate_tensor_shape,
@@ -93,6 +87,10 @@ def sample_mixture_of_diffusers(
sampler_name=sampler_name,
steps=steps,
denoise=denoise,
view=sampling_schedulers.SchedulerView(
latent_width=latent_tile_width,
latent_height=latent_tile_height,
),
).to(model.load_device)
latent_samples = validate_latent_samples(
@@ -217,8 +215,10 @@ class MixtureOfDiffusersModelWrapper:
self._plan = plan
self._existing_wrapper = existing_wrapper
self._tile_weights_2d: torch.Tensor | None = None
self._semantic_tile_weights = SemanticTileWeightCache(plan.tiles)
self._tile_predictions = TilePredictionAccumulator(
plan,
diffusion_mode="mixture_of_diffusers",
)
def __call__(
self,
@@ -248,35 +248,11 @@ class MixtureOfDiffusersModelWrapper:
"ControlNet in the first implementation."
)
output_buffer = torch.zeros_like(x)
weight_buffer = new_spatial_weight_buffer(x, self._plan)
input_batch_size = int(x.shape[0])
for batch in self._plan.batches:
tiled_args = self._make_tiled_args(
args=args,
tiles=batch,
input_batch_size=input_batch_size,
)
tile_output = self._call_original(apply_model, tiled_args)
weights = self._weights_for(tile_output)
accumulation_weights = weights.to(dtype=weight_buffer.dtype)
semantic_weights = self._semantic_tile_weights.for_output(tile_output)
for index, tile in enumerate(batch):
tile_slice = spatial_tile_slicer(tile, x.ndim)
start = index * input_batch_size
end = start + input_batch_size
model_weight, accumulation_weight = (
self._semantic_tile_weights.for_tile(
semantic_weights,
tile,
)
)
tile_weight = weights * model_weight
output_buffer[tile_slice] += tile_output[start:end] * tile_weight
weight_buffer[tile_slice] += accumulation_weights * accumulation_weight
return output_buffer / weight_buffer.to(dtype=output_buffer.dtype)
return self._tile_predictions.predict(
args=args,
x=x,
evaluate=lambda tiled_args: self._call_original(apply_model, tiled_args),
)
def _call_original(
self,
@@ -292,41 +268,6 @@ class MixtureOfDiffusersModelWrapper:
raise ValueError("Mixture of Diffusers conditioning must be a dict.")
return apply_model(args["input"], args["timestep"], **conditioning)
def _make_tiled_args(
self,
*,
args: dict[str, Any],
tiles: Sequence[LatentTile],
input_batch_size: int,
) -> dict[str, Any]:
"""Create apply-model args for one tile batch."""
return make_tiled_model_args(
args=args,
tiles=tiles,
input_batch_size=input_batch_size,
latent_height=self._plan.latent_height,
latent_width=self._plan.latent_width,
)
def _weights_for(self, x: torch.Tensor) -> torch.Tensor:
"""Return cached Gaussian tile weights for the active device and dtype."""
if (
self._tile_weights_2d is None
or self._tile_weights_2d.device != x.device
or self._tile_weights_2d.dtype != x.dtype
):
self._tile_weights_2d = gaussian_tile_weights(
self._plan.tile_width,
self._plan.tile_height,
device=x.device,
dtype=x.dtype,
)
return self._tile_weights_2d.reshape(
(1,) * (x.ndim - 2) + (self._plan.tile_height, self._plan.tile_width)
)
def _validate_supplied_plan(
plan: TiledDiffusionPlan,