sophia / modeling_sophia.py
Arain119
Sophia 1.0.0 — 1B K3-hybrid Chinese chat model (HF remote-code export + native package)
d53adc9
Raw History Blame
5.22 kB
# 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,
)