""" 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)