# Generated by ml.integrations.export.runtime_packager.write_remote_code_bundle. # Exported for HuggingFace trust_remote_code loading. # This file is intentionally self-contained. """ Sophia model implementation. This file contains: - `SophiaForCausalLM`: Hugging Face causal language model adapter over the Sophia runtime. """ from __future__ import annotations from typing import TYPE_CHECKING from .sophia_runtime import Transformer as _SophiaRuntimeTransformer # noqa: F401 if TYPE_CHECKING: from .sophia_runtime import * # noqa: F401,F403 from .model_config import * # noqa: F401,F403 from .hf_remote_code import * # noqa: F401,F403 from .hf_support import * # noqa: F401,F403 from .hf_cache import * # noqa: F401,F403 from .hf_config import * # noqa: F401,F403 from .hf_generation import * # noqa: F401,F403 from .hf_lifecycle import * # noqa: F401,F403 from .model_attention import * # noqa: F401,F403 from .model_blocks import * # noqa: F401,F403 from .model_ops import * # noqa: F401,F403 from .model_runtime import * # noqa: F401,F403 from .runtime_contracts import * # noqa: F401,F403 from .model_runtime_control import * # noqa: F401,F403 from .model_state import * # noqa: F401,F403 from .model_transformer_setup import * # noqa: F401,F403 from .runtime_linear import * # noqa: F401,F403 from .canonical_config import * # noqa: F401,F403 from .model_dir import * # noqa: F401,F403 from .loss_stats import * # noqa: F401,F403 from .input_mask import * # noqa: F401,F403 from .cache_decode import * # noqa: F401,F403 from .config_projection import * # noqa: F401,F403 from .hf_projection import * # noqa: F401,F403 from .runtime_backend import * # noqa: F401,F403 from .decoder_forward import * # noqa: F401,F403 from .decoder_types import * # noqa: F401,F403 from .decoder_loss import * # noqa: F401,F403 from .decoder_loss_forward import * # noqa: F401,F403 from .decoder_full import * # noqa: F401,F403 from .pretrained_bundle import * # noqa: F401,F403 from .decoder_output import * # noqa: F401,F403 from .decoder_runtime import * # noqa: F401,F403 from .decoder_host import * # noqa: F401,F403 from .sophia_decoder import * # noqa: F401,F403 from .semantics import * # noqa: F401,F403 import torch from transformers import PreTrainedModel from transformers.generation import GenerationMixin from transformers.modeling_outputs import CausalLMOutputWithPast from .runtime_backend import resolve_runtime_backend from .hf_cache import SophiaCache from .hf_generation import ( CausalLMForwardMixin, GenerationCacheMixin, ) from .hf_lifecycle import PreTrainedLifecycleMixin from .hf_projection import SophiaConfig from .decoder_output import ( format_hf_causal_lm_output as _format_hf_causal_lm_output, resolve_return_dict as _resolve_return_dict, ) from .decoder_host import DecoderHostMixin from .decoder_runtime import ( DecoderModelMixin, DecoderRecipeMixin, DecoderRuntimeMixin, ) class SophiaForCausalLM( DecoderModelMixin, DecoderRuntimeMixin, DecoderRecipeMixin, DecoderHostMixin, GenerationCacheMixin, CausalLMForwardMixin, PreTrainedLifecycleMixin, PreTrainedModel, GenerationMixin, ): config_class = SophiaConfig base_model_prefix = "model" _tied_weights_keys = {"model.output.weight": "model.tok_embeddings.weight"} supports_gradient_checkpointing = True _is_stateful = True @classmethod def _supports_default_dynamic_cache(cls) -> bool: return False def __init__( self, config: SophiaConfig | None = None, *, runtime_max_seq_len: int | None = None, ): config = config or SophiaConfig() super().__init__(config) self._initialize_decoder_runtime( config=config, runtime_max_seq_len=runtime_max_seq_len, runtime_backend=resolve_runtime_backend(), gradient_checkpointing_enabled=bool( getattr(self, "gradient_checkpointing", False) ), register_tied_weights=True, ) def _set_gradient_checkpointing( self, enable: bool = True, gradient_checkpointing_func: object = None, ) -> None: del gradient_checkpointing_func self.gradient_checkpointing = bool(enable) self._sync_runtime_gradient_checkpointing(enable=bool(enable)) def _resolve_runtime_return_dict(self, *, return_dict: bool | None) -> bool: return _resolve_return_dict( config=self.config, return_dict=return_dict, ) def _format_decoder_output( self, *, loss: torch.Tensor | None, logits: torch.Tensor | None, cache: SophiaCache | None, return_dict: bool, ) -> CausalLMOutputWithPast | tuple[torch.Tensor, ...]: return _format_hf_causal_lm_output( output_cls=CausalLMOutputWithPast, loss=loss, logits=logits, past_key_values=cache, return_dict=return_dict, )