"""Scoped XPO3 attention routing runtime for ComfyUI generation calls. This module adds one context-manager API that temporarily patches the packaged Mage-Flow Turbo attention route to use the validated official ``spas_sage2_attn_meansim_topk_cuda`` primitive under the bounded CFG1/four-step envelope. Importing this module does not import the Sparge dependency or touch CUDA. """ from __future__ import annotations from collections import Counter from contextlib import ExitStack, contextmanager from typing import Any, Callable, Iterator, Sequence, TypeVar EXPECTED_DENOISE_STEPS = 4 DEFAULT_SELECTED_STEPS = (1, 2) EXPECTED_HEADS = 24 EXPECTED_HEAD_DIM = 128 CANDIDATE_TOPK = 1.0 CANDIDATE_SMOOTH_K = True T = TypeVar("T") def _callable_identity(value: Any) -> tuple[Any, Any]: return ( getattr(value, "__func__", value), getattr(value, "__self__", None), ) @contextmanager def _temporary_attribute( target: Any, name: str, replacement: Any, ) -> Iterator[Any]: namespace = getattr(target, "__dict__", {}) had_instance_value = name in namespace original_instance_value = namespace.get(name) original_effective = getattr(target, name) setattr(target, name, replacement) try: yield original_effective finally: if had_instance_value: setattr(target, name, original_instance_value) else: delattr(target, name) def _candidate_route_enabled(gate_state: dict[str, Any]) -> bool: return bool(gate_state.get("step_enabled", False)) and bool( gate_state.get("block_enabled", False) ) def _normalize_int_set( values: Sequence[int] | set[int] | None, *, default: Sequence[int], ) -> set[int]: source = default if values is None else values result = {int(value) for value in source} if not result: raise ValueError("selected set must not be empty") return result def _report_template( *, enabled: bool, direct_hnd: bool, steps: int, static_shift: float, cfg: float, selected_steps: set[int], selected_blocks: set[int] | None, ) -> dict[str, Any]: return { "requested": { "enabled": bool(enabled), "direct_hnd": bool(direct_hnd), "steps": int(steps), "static_shift": float(static_shift), "cfg": float(cfg), "selected_steps": sorted(int(value) for value in selected_steps), "selected_blocks": ( None if selected_blocks is None else sorted(int(value) for value in selected_blocks) ), "topk": CANDIDATE_TOPK, "smooth_k": CANDIDATE_SMOOTH_K, "attention_backend": "official_spas_sage2_attn_meansim_topk_cuda", }, "active_feature": { "enabled": False, "patched": False, "mode": "fallback", "fallback_reason": None, "dependency_available": False, }, "routing": { "wrapper_calls": 0, "routed_calls": 0, "fallback_calls": 0, "routed_calls_by_step": {}, "routed_calls_by_block": {}, "fallback_calls_by_step": {}, "fallback_calls_by_block": {}, "route_records": [], }, "restoration": { "velocity_restored": None, "block_forwards_restored": None, "block_forward_instance_attribute_state_restored": None, "attention_callable_restored": None, "processors_restored": None, "processor_instance_attribute_state_restored": None, "all_restored": None, }, } class _RoutingReport: def __init__(self, report: dict[str, Any]) -> None: self.report = report self.calls_by_step: Counter[int] = Counter() self.calls_by_block: Counter[int] = Counter() self.fallback_by_step: Counter[int] = Counter() self.fallback_by_block: Counter[int] = Counter() @staticmethod def _index(state: dict[str, Any], key: str) -> int: value = state.get(key) return -1 if value is None else int(value) def record_fallback(self, gate_state: dict[str, Any]) -> None: routing = self.report["routing"] routing["wrapper_calls"] += 1 routing["fallback_calls"] += 1 self.fallback_by_step[self._index(gate_state, "step_index")] += 1 self.fallback_by_block[self._index(gate_state, "block_index")] += 1 self._flush() def record_route( self, gate_state: dict[str, Any], lengths: Sequence[int], ) -> None: routing = self.report["routing"] routing["wrapper_calls"] += 1 routing["routed_calls"] += 1 step_index = self._index(gate_state, "step_index") block_index = self._index(gate_state, "block_index") self.calls_by_step[step_index] += 1 self.calls_by_block[block_index] += 1 routing["route_records"].append( { "step_index": step_index, "block_index": block_index, "segment_lengths": [int(value) for value in lengths], } ) self._flush() def _flush(self) -> None: routing = self.report["routing"] routing["routed_calls_by_step"] = { str(key): value for key, value in sorted(self.calls_by_step.items()) } routing["routed_calls_by_block"] = { str(key): value for key, value in sorted(self.calls_by_block.items()) } routing["fallback_calls_by_step"] = { str(key): value for key, value in sorted(self.fallback_by_step.items()) } routing["fallback_calls_by_block"] = { str(key): value for key, value in sorted(self.fallback_by_block.items()) } def _cumulative_lengths_to_list( values: Any, *, torch: Any, ) -> list[int] | None: if values is None: return None if isinstance(values, torch.Tensor): if values.ndim != 1: return None result = [int(value) for value in values.tolist()] else: try: result = [int(value) for value in values] except TypeError: return None if len(result) < 2 or result[0] != 0: return None if any(right <= left for left, right in zip(result[:-1], result[1:])): return None return result def _validate_wrapper_candidate_call( q: Any, k: Any, v: Any, *, cu_q: list[int] | None, cu_k: list[int] | None, dropout_p: float, causal: bool, window_size: tuple[int | None, int | None], softcap: float, alibi_slopes: Any, deterministic: bool, return_attn_probs: bool, block_table: Any, extra_kwargs: dict[str, Any], max_seqlen_q: int | None, max_seqlen_k: int | None, torch: Any, ) -> tuple[bool, list[int]]: if extra_kwargs: return False, [] if cu_q is None or cu_k is None: return False, [] if q.ndim != 3 or k.ndim != 3 or v.ndim != 3: return False, [] if q.shape != k.shape or k.shape != v.shape: return False, [] if q.dtype != torch.bfloat16 or k.dtype != torch.bfloat16 or v.dtype != torch.bfloat16: return False, [] if q.device != k.device or k.device != v.device: return False, [] if int(q.shape[1]) != EXPECTED_HEADS or int(q.shape[2]) != EXPECTED_HEAD_DIM: return False, [] if len(cu_q) != len(cu_k) or cu_q[-1] != int(q.shape[0]) or cu_k[-1] != int(k.shape[0]): return False, [] segment_lengths = [] for qs, qe, ks, ke in zip(cu_q[:-1], cu_q[1:], cu_k[:-1], cu_k[1:]): q_len = int(qe - qs) k_len = int(ke - ks) if q_len != k_len or q_len <= 0: return False, [] segment_lengths.append(q_len) if not segment_lengths or min(segment_lengths) < 128: return False, [] if max_seqlen_q is not None and int(max_seqlen_q) != max(segment_lengths): return False, [] if max_seqlen_k is not None and int(max_seqlen_k) != max(segment_lengths): return False, [] if float(dropout_p) != 0.0: return False, [] if bool(causal): return False, [] if window_size not in ((-1, -1), (None, None)): return False, [] if float(softcap) != 0.0: return False, [] if alibi_slopes is not None: return False, [] if bool(deterministic): return False, [] if bool(return_attn_probs): return False, [] if block_table is not None: return False, [] return True, segment_lengths def _nhd_to_hnd(segment: Any) -> Any: return segment.permute(1, 0, 2).unsqueeze(0).contiguous() def _hnd_to_nhd(segment: Any) -> Any: return segment.squeeze(0).permute(1, 0, 2).contiguous() def _make_sparge_wrapper( *, gate_state: dict[str, Any], flash_fallback: Callable[..., Any], routing: _RoutingReport, sparge_fn: Callable[..., Any], torch: Any, ) -> Callable[..., Any]: def wrapped_flash_attn_varlen_func( q: Any, k: Any, v: Any, cu_seqlens_q: Any = None, cu_seqlens_k: Any = None, max_seqlen_q: int | None = None, max_seqlen_k: int | None = None, dropout_p: float = 0.0, softmax_scale: float | None = None, causal: bool = False, window_size: tuple[int | None, int | None] = (-1, -1), softcap: float = 0.0, alibi_slopes: Any = None, deterministic: bool = False, return_attn_probs: bool = False, block_table: Any = None, **extra_kwargs: Any, ) -> Any: if not _candidate_route_enabled(gate_state): routing.record_fallback(gate_state) return flash_fallback( q, k, v, cu_seqlens_q=cu_seqlens_q, cu_seqlens_k=cu_seqlens_k, max_seqlen_q=max_seqlen_q, max_seqlen_k=max_seqlen_k, dropout_p=dropout_p, softmax_scale=softmax_scale, causal=causal, window_size=window_size, softcap=softcap, alibi_slopes=alibi_slopes, deterministic=deterministic, return_attn_probs=return_attn_probs, block_table=block_table, **extra_kwargs, ) cu_q = _cumulative_lengths_to_list(cu_seqlens_q, torch=torch) cu_k = _cumulative_lengths_to_list(cu_seqlens_k, torch=torch) supported, segment_lengths = _validate_wrapper_candidate_call( q, k, v, cu_q=cu_q, cu_k=cu_k, dropout_p=dropout_p, causal=causal, window_size=window_size, softcap=softcap, alibi_slopes=alibi_slopes, deterministic=deterministic, return_attn_probs=return_attn_probs, block_table=block_table, extra_kwargs=extra_kwargs, max_seqlen_q=max_seqlen_q, max_seqlen_k=max_seqlen_k, torch=torch, ) if not supported: routing.record_fallback(gate_state) return flash_fallback( q, k, v, cu_seqlens_q=cu_seqlens_q, cu_seqlens_k=cu_seqlens_k, max_seqlen_q=max_seqlen_q, max_seqlen_k=max_seqlen_k, dropout_p=dropout_p, softmax_scale=softmax_scale, causal=causal, window_size=window_size, softcap=softcap, alibi_slopes=alibi_slopes, deterministic=deterministic, return_attn_probs=return_attn_probs, block_table=block_table, **extra_kwargs, ) routing.record_route(gate_state, segment_lengths) outputs_hnd = [] for qs, qe, ks, ke in zip(cu_q[:-1], cu_q[1:], cu_k[:-1], cu_k[1:]): output_hnd = sparge_fn( _nhd_to_hnd(q[qs:qe]), _nhd_to_hnd(k[ks:ke]), _nhd_to_hnd(v[ks:ke]), attn_mask=None, dropout_p=0.0, is_causal=False, scale=softmax_scale, smooth_k=CANDIDATE_SMOOTH_K, topk=CANDIDATE_TOPK, tensor_layout="HND", return_sparsity=False, ) if isinstance(output_hnd, tuple): output_hnd = output_hnd[0] outputs_hnd.append(output_hnd) return torch.cat([_hnd_to_nhd(output) for output in outputs_hnd], dim=0) return wrapped_flash_attn_varlen_func class _DirectSingleSampleSpargeProcessor: def __init__( self, *, original: Any, gate_state: dict[str, Any], routing: _RoutingReport, sparge_fn: Callable[..., Any], torch: Any, ) -> None: self.original = original self.gate_state = gate_state self.routing = routing self.sparge_fn = sparge_fn self.torch = torch def __call__( self, attn: Any, hidden_states: Any, img_cu_lens: Any, attention_mask: Any = None, encoder_hidden_states: Any = None, txt_cu_lens: Any = None, image_rotary_emb: Any = None, **kwargs: Any, ) -> Any: if not _candidate_route_enabled(self.gate_state): self.routing.record_fallback(self.gate_state) return self.original( attn, hidden_states, img_cu_lens, attention_mask=attention_mask, encoder_hidden_states=encoder_hidden_states, txt_cu_lens=txt_cu_lens, image_rotary_emb=image_rotary_emb, **kwargs, ) txt_cu = _cumulative_lengths_to_list(txt_cu_lens, torch=self.torch) img_cu = _cumulative_lengths_to_list(img_cu_lens, torch=self.torch) supported = ( encoder_hidden_states is not None and attention_mask is None and image_rotary_emb is not None and hidden_states.ndim == 3 and encoder_hidden_states.ndim == 3 and hidden_states.shape[0] == 1 and encoder_hidden_states.shape[0] == 1 and txt_cu is not None and img_cu is not None and len(txt_cu) == 2 and len(img_cu) == 2 and hidden_states.dtype == self.torch.bfloat16 and encoder_hidden_states.dtype == self.torch.bfloat16 and hidden_states.device == encoder_hidden_states.device and int(getattr(attn, "heads", -1)) == EXPECTED_HEADS and not kwargs ) txt_tokens = -1 if txt_cu is None else int(txt_cu[-1]) img_tokens = -1 if img_cu is None else int(img_cu[-1]) if supported: supported = ( txt_tokens == int(encoder_hidden_states.shape[1]) and img_tokens == int(hidden_states.shape[1]) and (txt_tokens + img_tokens) >= 128 ) if not supported: self.routing.record_fallback(self.gate_state) return self.original( attn, hidden_states, img_cu_lens, attention_mask=attention_mask, encoder_hidden_states=encoder_hidden_states, txt_cu_lens=txt_cu_lens, image_rotary_emb=image_rotary_emb, **kwargs, ) from mage_flow.models.modules.mage_layers import apply_rotary_emb_mageflow if getattr(attn, "to_qkv", None) is not None: img_query, img_key, img_value = attn.to_qkv(hidden_states).chunk(3, dim=-1) else: img_query = attn.to_q(hidden_states) img_key = attn.to_k(hidden_states) img_value = attn.to_v(hidden_states) if getattr(attn, "add_qkv_proj", None) is not None: txt_query, txt_key, txt_value = attn.add_qkv_proj( encoder_hidden_states ).chunk(3, dim=-1) else: txt_query = attn.add_q_proj(encoder_hidden_states) txt_key = attn.add_k_proj(encoder_hidden_states) txt_value = attn.add_v_proj(encoder_hidden_states) img_query = img_query.unflatten(-1, (attn.heads, -1)).flatten(0, 1) img_key = img_key.unflatten(-1, (attn.heads, -1)).flatten(0, 1) img_value = img_value.unflatten(-1, (attn.heads, -1)).flatten(0, 1) txt_query = txt_query.unflatten(-1, (attn.heads, -1)).flatten(0, 1) txt_key = txt_key.unflatten(-1, (attn.heads, -1)).flatten(0, 1) txt_value = txt_value.unflatten(-1, (attn.heads, -1)).flatten(0, 1) expected_img_shape = (img_tokens, EXPECTED_HEADS, EXPECTED_HEAD_DIM) expected_txt_shape = (txt_tokens, EXPECTED_HEADS, EXPECTED_HEAD_DIM) if any( tuple(tensor.shape) != expected_shape for tensor, expected_shape in ( (img_query, expected_img_shape), (img_key, expected_img_shape), (img_value, expected_img_shape), (txt_query, expected_txt_shape), (txt_key, expected_txt_shape), (txt_value, expected_txt_shape), ) ): self.routing.record_fallback(self.gate_state) return self.original( attn, hidden_states, img_cu_lens, attention_mask=attention_mask, encoder_hidden_states=encoder_hidden_states, txt_cu_lens=txt_cu_lens, image_rotary_emb=image_rotary_emb, **kwargs, ) if attn.norm_q is not None: img_query = attn.norm_q(img_query) if attn.norm_k is not None: img_key = attn.norm_k(img_key) if attn.norm_added_q is not None: txt_query = attn.norm_added_q(txt_query) if attn.norm_added_k is not None: txt_key = attn.norm_added_k(txt_key) img_query = apply_rotary_emb_mageflow(img_query, image_rotary_emb) img_key = apply_rotary_emb_mageflow(img_key, image_rotary_emb) def pack_joint_hnd(txt_tensor: Any, img_tensor: Any) -> Any: return self.torch.cat( (txt_tensor.transpose(0, 1), img_tensor.transpose(0, 1)), dim=1, ).unsqueeze(0) joint_query = pack_joint_hnd(txt_query, img_query) joint_key = pack_joint_hnd(txt_key, img_key) joint_value = pack_joint_hnd(txt_value, img_value) self.routing.record_route(self.gate_state, [txt_tokens + img_tokens]) joint_attn_output = self.sparge_fn( joint_query, joint_key, joint_value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None, smooth_k=CANDIDATE_SMOOTH_K, topk=CANDIDATE_TOPK, tensor_layout="HND", return_sparsity=False, ) if isinstance(joint_attn_output, tuple): joint_attn_output = joint_attn_output[0] joint_bshd = joint_attn_output.transpose(1, 2) txt_attn_output = joint_bshd[:, :txt_tokens].reshape( txt_tokens, attn.heads * EXPECTED_HEAD_DIM, ) img_attn_output = joint_bshd[:, txt_tokens:].reshape( img_tokens, attn.heads * EXPECTED_HEAD_DIM, ) img_attn_output = img_attn_output.to(txt_query.dtype) txt_attn_output = txt_attn_output.to(txt_query.dtype) img_attn_output = attn.to_out[0](img_attn_output) if len(attn.to_out) > 1: img_attn_output = attn.to_out[1](img_attn_output) txt_attn_output = attn.to_add_out(txt_attn_output) txt_attn_output = txt_attn_output.view( encoder_hidden_states.shape[0], encoder_hidden_states.shape[1], txt_attn_output.shape[-1], ) return img_attn_output, txt_attn_output def _build_sigma_to_step_index( *, model: Any, steps: int, static_shift: float, torch: Any, ) -> dict[float, int]: import mage_flow.pipeline as mage_pipeline scheduler = mage_pipeline._get_scheduler( model, int(steps), torch.device("cuda:0"), float(static_shift), ) return { round(float(sigma.item()), 8): index for index, sigma in enumerate(scheduler.sigmas) } def _load_sparge_dependency() -> Callable[..., Any]: from spas_sage_attn import spas_sage2_attn_meansim_topk_cuda return spas_sage2_attn_meansim_topk_cuda @contextmanager def _patch_velocity_gate( *, allowed_steps: set[int], step_index_by_sigma: dict[float, int], gate_state: dict[str, Any], ) -> Iterator[None]: import mage_flow.pipeline as mage_pipeline original_velocity = mage_pipeline._velocity def wrapped_velocity( transformer: Any, image: Any, context: dict[str, Any], sigma: float, ) -> Any: sigma_key = round(float(sigma), 8) step_index = step_index_by_sigma.get(sigma_key) if step_index is None: return original_velocity(transformer, image, context, sigma) previous_step = gate_state.get("step_index") previous_enabled = gate_state.get("step_enabled", False) gate_state["step_index"] = step_index gate_state["step_enabled"] = step_index in allowed_steps try: return original_velocity(transformer, image, context, sigma) finally: gate_state["step_index"] = previous_step gate_state["step_enabled"] = previous_enabled try: with _temporary_attribute(mage_pipeline, "_velocity", wrapped_velocity): yield finally: gate_state["step_enabled"] = False gate_state["step_index"] = None @contextmanager def _patch_selected_transformer_blocks( *, transformer: Any, selected_blocks: set[int], gate_state: dict[str, Any], ) -> Iterator[None]: blocks = list(transformer.transformer_blocks) if any(index < 0 or index >= len(blocks) for index in selected_blocks): raise ValueError("selected block index is out of range") def wrap_forward( original_forward: Callable[..., Any], block_index: int, ) -> Callable[..., Any]: def wrapped_forward(*args: Any, **kwargs: Any) -> Any: previous_block = gate_state.get("block_index") previous_enabled = gate_state.get("block_enabled", False) gate_state["block_index"] = block_index gate_state["block_enabled"] = True try: return original_forward(*args, **kwargs) finally: gate_state["block_index"] = previous_block gate_state["block_enabled"] = previous_enabled return wrapped_forward with ExitStack() as stack: for block_index in sorted(selected_blocks): block = blocks[block_index] stack.enter_context( _temporary_attribute( block, "forward", wrap_forward(block.forward, block_index), ) ) try: yield finally: gate_state["block_enabled"] = False gate_state["block_index"] = None def _mark_fallback(report: dict[str, Any], reason: str) -> None: report["active_feature"]["enabled"] = False report["active_feature"]["patched"] = False report["active_feature"]["mode"] = "fallback" report["active_feature"]["fallback_reason"] = reason restoration = report["restoration"] for key in restoration: if restoration[key] is None: restoration[key] = "not_applicable" @contextmanager def xpo3_attention_runtime( *, pipe: Any, torch: Any, enabled: bool, direct_hnd: bool, steps: int, static_shift: float, cfg: float, selected_steps: Sequence[int] | set[int] | None = None, selected_blocks: Sequence[int] | set[int] | None = None, required_cfg: float = 1.0, expected_steps: int = EXPECTED_DENOISE_STEPS, ) -> Iterator[dict[str, Any]]: """Temporarily enable the validated XPO3 Sparge/Sage2 attention route. Yields a mutable report dict describing whether the route was patched, whether it fell back exactly to the original path, live routing counters, and post-context restoration status. """ normalized_steps = _normalize_int_set( selected_steps, default=DEFAULT_SELECTED_STEPS, ) normalized_blocks = ( None if selected_blocks is None else _normalize_int_set(selected_blocks, default=()) ) report = _report_template( enabled=enabled, direct_hnd=direct_hnd, steps=steps, static_shift=static_shift, cfg=cfg, selected_steps=normalized_steps, selected_blocks=normalized_blocks, ) report["requested"]["required_cfg"] = float(required_cfg) report["requested"]["expected_steps"] = int(expected_steps) if not enabled: _mark_fallback(report, "disabled") yield report return if float(cfg) != float(required_cfg): reason = "cfg_not_1" if float(required_cfg) == 1.0 else "cfg_not_allowed" _mark_fallback(report, reason) yield report return if int(steps) != int(expected_steps): reason = ( "steps_not_4" if int(expected_steps) == EXPECTED_DENOISE_STEPS else "steps_not_expected" ) _mark_fallback(report, reason) yield report return transformer = getattr(getattr(pipe, "model", None), "transformer", None) blocks = list(getattr(transformer, "transformer_blocks", [])) if transformer is not None else [] if transformer is None or not blocks: _mark_fallback(report, "unsupported_pipe") yield report return active_blocks = ( set(range(len(blocks))) if normalized_blocks is None else set(normalized_blocks) ) report["requested"]["selected_blocks"] = sorted(active_blocks) if any(index < 0 or index >= len(blocks) for index in active_blocks): raise ValueError("selected_blocks contains an out-of-range block index") try: sparge_fn = _load_sparge_dependency() report["active_feature"]["dependency_available"] = True except Exception as exc: report["active_feature"]["dependency_available"] = False report["active_feature"]["dependency_error"] = ( f"{type(exc).__name__}: {exc}" ) _mark_fallback(report, "dependency_missing") yield report return import mage_flow.models.modules.mage_layers as mage_layers import mage_flow.pipeline as mage_pipeline step_index_by_sigma = _build_sigma_to_step_index( model=pipe.model, steps=steps, static_shift=static_shift, torch=torch, ) report["active_feature"]["enabled"] = True report["active_feature"]["patched"] = True report["active_feature"]["mode"] = "direct_hnd" if direct_hnd else "wrapper" report["active_feature"]["fallback_reason"] = None report["active_feature"]["selected_steps"] = sorted(normalized_steps) report["active_feature"]["selected_blocks"] = sorted(active_blocks) report["active_feature"]["step_index_by_sigma"] = dict(step_index_by_sigma) routing = _RoutingReport(report) gate_state: dict[str, Any] = { "step_enabled": False, "block_enabled": False, "step_index": None, "block_index": None, } original_velocity = _callable_identity(mage_pipeline._velocity) original_attention = _callable_identity(mage_layers.flash_attn_varlen_func) original_block_forwards = [_callable_identity(block.forward) for block in blocks] original_block_instance_flags = [ "forward" in getattr(block, "__dict__", {}) for block in blocks ] original_processors = [ _callable_identity(block.attn.processor) for block in blocks ] original_processor_instance_flags = [ "processor" in getattr(block.attn, "__dict__", {}) for block in blocks ] try: with ExitStack() as stack: stack.enter_context( _patch_velocity_gate( allowed_steps=normalized_steps, step_index_by_sigma=step_index_by_sigma, gate_state=gate_state, ) ) stack.enter_context( _patch_selected_transformer_blocks( transformer=transformer, selected_blocks=active_blocks, gate_state=gate_state, ) ) if direct_hnd: for block_index in sorted(active_blocks): attn = blocks[block_index].attn replacement = _DirectSingleSampleSpargeProcessor( original=attn.processor, gate_state=gate_state, routing=routing, sparge_fn=sparge_fn, torch=torch, ) stack.enter_context( _temporary_attribute(attn, "processor", replacement) ) else: wrapper = _make_sparge_wrapper( gate_state=gate_state, flash_fallback=mage_layers.flash_attn_varlen_func, routing=routing, sparge_fn=sparge_fn, torch=torch, ) stack.enter_context( _temporary_attribute( mage_layers, "flash_attn_varlen_func", wrapper, ) ) yield report finally: restoration = report["restoration"] restoration["velocity_restored"] = ( _callable_identity(mage_pipeline._velocity) == original_velocity ) restoration["block_forwards_restored"] = all( _callable_identity(block.forward) == original for block, original in zip(blocks, original_block_forwards) ) restoration["block_forward_instance_attribute_state_restored"] = all( ("forward" in getattr(block, "__dict__", {})) == original_flag for block, original_flag in zip(blocks, original_block_instance_flags) ) restoration["attention_callable_restored"] = ( "not_applicable" if direct_hnd else _callable_identity(mage_layers.flash_attn_varlen_func) == original_attention ) restoration["processors_restored"] = ( all( _callable_identity(block.attn.processor) == original for block, original in zip(blocks, original_processors) ) if direct_hnd else "not_applicable" ) restoration["processor_instance_attribute_state_restored"] = ( all( ("processor" in getattr(block.attn, "__dict__", {})) == original_flag for block, original_flag in zip( blocks, original_processor_instance_flags, ) ) if direct_hnd else "not_applicable" ) restoration_checks = [] for key, value in restoration.items(): if key == "all_restored": continue if value == "not_applicable": continue restoration_checks.append(bool(value)) restoration["all_restored"] = all(restoration_checks) __all__ = ["xpo3_attention_runtime"]