Generalise number of denoising steps

This commit is contained in:
ciaranbor
2026-01-06 10:51:20 +00:00
parent 129df1ec89
commit e8b0a2124c
3 changed files with 13 additions and 10 deletions
@@ -34,11 +34,17 @@ class DistributedFlux1:
(group, rank, world_size, shard boundaries).
"""
__slots__ = ("_model", "_group", "_shard_metadata")
__slots__ = ("_model", "_group", "_shard_metadata", "_variant")
_model: Flux1
_group: Optional[mx.distributed.Group]
_shard_metadata: PipelineShardMetadata
_variant: str
_steps_map: dict[str, dict[str, int]] = {
"schnell": {"low": 1, "medium": 2, "high": 4},
"dev": {"low": 10, "medium": 25, "high": 50},
}
def __init__(
self,
@@ -58,6 +64,7 @@ class DistributedFlux1:
group: MLX distributed group for multi-node coordination (None for single-node)
quantize: Optional quantization bit width
"""
variant = model_id.split("-")[-1]
model = Flux1(
model_config=ModelConfig.from_name(model_name=model_id, base_model=None),
local_path=str(local_path),
@@ -87,6 +94,7 @@ class DistributedFlux1:
object.__setattr__(self, "_model", model)
object.__setattr__(self, "_group", group)
object.__setattr__(self, "_shard_metadata", shard_metadata)
object.__setattr__(self, "_variant", variant)
@classmethod
def from_bound_instance(cls, bound_instance: BoundInstance) -> "DistributedFlux1":
@@ -205,11 +213,7 @@ class DistributedFlux1:
Generated PIL Image (rank 0) or None (other ranks)
"""
# Determine number of inference steps based on quality
steps = 2
if quality == "low":
steps = 1
elif quality == "high":
steps = 4
steps = self._steps_map[self._variant][quality]
config = Config(num_inference_steps=steps, height=height, width=width)
image = self._generate_image(settings=config, prompt=prompt, seed=seed)
@@ -58,7 +58,7 @@ def mflux_generate(
# Parse parameters
width, height = parse_size(task.size)
quality: Literal["low", "medium", "high"] = task.quality or "medium"
seed = 2 # TODO: not in OAI API?
seed = 2 # TODO: not in OAI API? random?
# Generate using the model's generate method
image = model.generate(
@@ -204,12 +204,11 @@ class DistributedDenoising:
hidden_states, self.rank + 1, group=self.group
)
mx.eval(hidden_states)
mx_barrier(group=self.group)
#
# === PHASE 5: All-gather Final Output ===
# All stages participate to receive the final output
mx.eval(hidden_states)
mx_barrier(group=self.group)
hidden_states = mx.distributed.all_gather(hidden_states, group=self.group)[
-hidden_states.shape[0] :
]