Sky-v2.0-Lite / modeling_sky_crest.py
Atharvsinh's picture
Upload folder using huggingface_hub
46cc6c9 verified
Raw
History Blame Contribute Delete
1.4 kB
"""
Sky CREST Model — 0labs
SkyCRESTForCausalLM: Adaptive-depth language model architecture.
"""
import torch
import torch.nn as nn
import transformers
from .configuration_sky_crest import SkyCRESTConfig
from .crest_block import CRESTBlock
# Dynamically resolve base architecture
_BASE_CLASSES = [
"Qwen3_5ForCausalLM",
"Qwen2ForCausalLM",
"LlamaForCausalLM",
]
_BaseClass = None
for _name in _BASE_CLASSES:
_BaseClass = getattr(transformers, _name, None)
if _BaseClass is not None:
break
if _BaseClass is None:
raise ImportError(
"Sky v2.0 requires transformers>=4.51.0. "
"Run: pip install --upgrade transformers"
)
class SkyCRESTForCausalLM(_BaseClass):
"""Sky v2.0 — Adaptive-depth language model with CREST architecture by 0labs."""
config_class = SkyCRESTConfig
def __init__(self, config):
super().__init__(config)
max_steps = getattr(config, "crest_max_steps", 4)
hidden_size = config.hidden_size
# Find the layers — handle both model.layers and model.language_model.layers
if hasattr(self.model, "language_model"):
layers = self.model.language_model.layers
else:
layers = self.model.layers
for layer in layers:
orig_mlp = layer.mlp
layer.mlp = CRESTBlock(orig_mlp, hidden_size=hidden_size, max_steps=max_steps)