multimodalart's picture
multimodalart HF Staff
Bernini-Diffusers-v2 r2v demo
fed6c68 verified
Raw
History Blame Contribute Delete
4.92 kB
# Copyright 2025 Bytedance Ltd. and/or its affiliates
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from __future__ import annotations
from typing import TYPE_CHECKING
from ..utils import logging
from ..utils.env import get_env
# Eagerly import kernel packages so that every op registers itself with the
# registry. Order does not matter; each ``register_op`` call is idempotent.
from . import kernels, liger # noqa: F401 triggers all register_op() calls
from .config.registry import apply_global_ops
from .config.singleton import set_ops_config
from .dispatch import OpSlot
from .kernels import attention, cross_entropy, load_balancing_loss, moe # noqa: F401
from .kernels.load_balancing_loss import load_balancing_loss_func
from .kernels.moe import fused_moe_forward
if TYPE_CHECKING:
from ..arguments.arguments_types import OpsImplementationConfig
__all__ = [
"fused_moe_forward",
"OpSlot",
"load_balancing_loss_func",
]
logger = logging.get_logger(__name__)
def build_ALL_OPS():
return [
("_fused_moe_forward", moe._fused_moe_forward),
("_flash_attention_forward", attention._flash_attention_forward),
("_load_balancing_loss", load_balancing_loss._load_balancing_loss),
]
def apply_ops_patch():
"""Import-time ops patch — attention only.
Registers VeOmni's SP-aware attention variants into the shared
``ALL_ATTENTION_FUNCTIONS`` registry. Loss dispatch (``LOSS_MAPPING``) is
deferred to ``apply_ops_config`` so there is a single binding point that
consumes ``OpsImplementationConfig``; ``build_foundation_model`` invokes
it automatically when callers pass ``ops_implementation=...`` (and
installs defaults otherwise).
"""
modeling_backend = get_env("MODELING_BACKEND")
if modeling_backend == "hf":
logger.info_rank0("⚠️ Skip applying ops patch. Using huggingface transformers backend.")
else:
from .kernels.attention import apply_veomni_attention_patch
apply_veomni_attention_patch()
logger.info_rank0("✅ VeOmni attention patches applied.")
def apply_ops_config(ops_config: OpsImplementationConfig) -> None:
"""Apply kernel patches based on resolved ``OpsImplementationConfig``.
Single install point for config-driven dispatch:
1. Binds the cross-entropy kernel into ``LOSS_MAPPING`` via
``install_loss_mapping`` (pre-bound ``partial`` — no runtime resolution).
2. Walks GLOBAL ops (e.g. load-balancing loss) and binds each selected
backend to its ``global_slot``.
3. Populates the ops-config singleton so per-model ``device_patch.py`` and
``OpSlot.bind`` can read the user's selections.
MoE dispatch is applied in ``build_foundation_model`` (via
``moe_implementation`` ∈ {``eager``, ``fused_triton``, ``fused_quack``,
``fused_npu``}); per-model kernels are applied by each model's
``device_patch.py``.
"""
set_ops_config(ops_config)
modeling_backend = get_env("MODELING_BACKEND")
if modeling_backend == "hf":
return
from .kernels.cross_entropy import install_loss_mapping
ce_label = install_loss_mapping(ops_config.cross_entropy_loss_implementation)
applied = apply_global_ops(ops_config)
applied.insert(0, ce_label)
logger.info_rank0(f"✅ VeOmni ops config applied: {', '.join(applied)}.")
logger.info_rank0(format_kernel_functions())
def format_kernel_functions() -> str:
lines = []
lines.append("\n=========== OPS ============")
for alias, func in build_ALL_OPS():
impl = func.__name__ if func is not None else "None"
lines.append(f"{alias} = {impl}")
# Cross-entropy is bound via LOSS_MAPPING (partial-wrapped), not a module
# global — surface it here so the log still shows the active CE kernel.
lines.append(f"cross_entropy = {_current_cross_entropy_name()}")
lines.append("==============================")
return "\n".join(lines)
def _current_cross_entropy_name() -> str:
from functools import partial
from transformers.loss.loss_utils import LOSS_MAPPING
entry = LOSS_MAPPING.get("ForCausalLM")
if entry is None:
return "unset"
if isinstance(entry, partial):
ce_fn = entry.keywords.get("cross_entropy_fn")
return getattr(ce_fn, "__name__", repr(ce_fn)) if ce_fn is not None else "unset"
return getattr(entry, "__name__", repr(entry))