"""Custom loader for the compressed Qwen3 checkpoint. A standard Qwen3 with a reduced number of layers, plus two per-layer buffers (`recover_scale`, `recover_bias`) applied at the start of each decoder layer's forward. Layers carry scale=1, bias=0 where no correction is present. Load with: AutoModelForCausalLM.from_pretrained(path, trust_remote_code=True) """ import torch import torch.nn as nn from transformers.models.qwen3.configuration_qwen3 import Qwen3Config from transformers.models.qwen3.modeling_qwen3 import ( Qwen3ForCausalLM, Qwen3Model, Qwen3DecoderLayer) class Qwen3RecoveredConfig(Qwen3Config): model_type = "qwen3_recovered" class RecoveredDecoderLayer(Qwen3DecoderLayer): def __init__(self, config, layer_idx): super().__init__(config, layer_idx) h = config.hidden_size self.register_buffer("recover_scale", torch.ones(h), persistent=True) self.register_buffer("recover_bias", torch.zeros(h), persistent=True) def forward(self, hidden_states, *args, **kwargs): s = self.recover_scale.to(hidden_states.dtype) b = self.recover_bias.to(hidden_states.dtype) hidden_states = hidden_states * s + b return super().forward(hidden_states, *args, **kwargs) class Qwen3RecoveredModel(Qwen3Model): config_class = Qwen3RecoveredConfig def __init__(self, config): super().__init__(config) self.layers = nn.ModuleList( [RecoveredDecoderLayer(config, i) for i in range(config.num_hidden_layers)]) self.post_init() class Qwen3RecoveredForCausalLM(Qwen3ForCausalLM): config_class = Qwen3RecoveredConfig def __init__(self, config): super().__init__(config) self.model = Qwen3RecoveredModel(config) self.post_init()