Qwen3-11B-30pct-Compressed-14B-EN-V1 / modeling_qwen3_recovered.py
Vincent-Daniel Yun
Qwen3-11B (30% compressed from Qwen3-14B) — E-AI Project
d390fc1 verified
Raw
History Blame Contribute Delete
1.79 kB
"""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()