# modeling_shared_reconstructor.py (auto-generated) import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from .configuration_shared_reconstructor import SharedReconstructorConfig class FiLMLayerNorm(nn.Module): def __init__(self, hidden_size, condition_size, eps=1e-5): super().__init__() self.norm = nn.LayerNorm(hidden_size, eps=eps, elementwise_affine=True) self.scale_shift = nn.Linear(condition_size, hidden_size * 2) nn.init.zeros_(self.scale_shift.weight) nn.init.zeros_(self.scale_shift.bias) def forward(self, x, condition): scale, shift = self.scale_shift(condition).chunk(2, dim=-1) normalized = self.norm(x) return normalized * (1.0 + scale[:, None, :]) + shift[:, None, :] class SharedLayerReconstructor(nn.Module): def __init__(self, hidden_size, num_teacher_layers, num_heads, ffn_size, layer_embedding_size, dropout, layer_norm_eps=1e-5): super().__init__() if hidden_size % num_heads != 0: raise ValueError("hidden_size must be divisible by num_heads.") self.hidden_size = hidden_size self.num_teacher_layers = num_teacher_layers self.num_heads = num_heads self.head_dim = hidden_size // num_heads self.dropout_rate = dropout self.layer_embedding = nn.Embedding(num_teacher_layers, layer_embedding_size) self.norm1 = FiLMLayerNorm(hidden_size, layer_embedding_size, eps=layer_norm_eps) self.q_proj = nn.Linear(hidden_size, hidden_size, bias=True) self.k_proj = nn.Linear(hidden_size, hidden_size, bias=True) self.v_proj = nn.Linear(hidden_size, hidden_size, bias=True) self.attention_output = nn.Linear(hidden_size, hidden_size, bias=True) self.norm2 = FiLMLayerNorm(hidden_size, layer_embedding_size, eps=layer_norm_eps) self.fc = nn.Linear(hidden_size, ffn_size, bias=True) self.proj = nn.Linear(ffn_size, hidden_size, bias=True) self.residual_gates = nn.Embedding(num_teacher_layers, 2) nn.init.constant_(self.residual_gates.weight, -2.0) self.dropout = nn.Dropout(dropout) def causal_attention(self, x): batch_size, sequence_length, _ = x.shape query = self.q_proj(x).view(batch_size, sequence_length, self.num_heads, self.head_dim).transpose(1, 2) key = self.k_proj(x).view(batch_size, sequence_length, self.num_heads, self.head_dim).transpose(1, 2) value = self.v_proj(x).view(batch_size, sequence_length, self.num_heads, self.head_dim).transpose(1, 2) attended = F.scaled_dot_product_attention( query, key, value, attn_mask=None, dropout_p=self.dropout_rate if self.training else 0.0, is_causal=True) attended = attended.transpose(1, 2).contiguous().view( batch_size, sequence_length, self.hidden_size) return self.attention_output(attended) def forward(self, hidden_states, layer_ids): condition = self.layer_embedding(layer_ids) residual_gates = torch.sigmoid(self.residual_gates(layer_ids)) normalized = self.norm1(hidden_states, condition) attention_output = self.causal_attention(normalized) hidden_states = hidden_states + residual_gates[:, 0, None, None] * self.dropout(attention_output) normalized = self.norm2(hidden_states, condition) ffn_output = self.proj(F.gelu(self.fc(normalized))) hidden_states = hidden_states + residual_gates[:, 1, None, None] * self.dropout(ffn_output) return hidden_states def rollout(self, hidden_states, start_layer=0, end_layer=None): if end_layer is None: end_layer = self.num_teacher_layers batch_size = hidden_states.shape[0] for layer_index in range(start_layer, end_layer): layer_ids = torch.full((batch_size,), layer_index, dtype=torch.long, device=hidden_states.device) hidden_states = self(hidden_states, layer_ids) return hidden_states class HFSharedReconstructor(PreTrainedModel): config_class = SharedReconstructorConfig base_model_prefix = "reconstructor" main_input_name = "hidden_states" def __init__(self, config): super().__init__(config) self.reconstructor = SharedLayerReconstructor( hidden_size=config.hidden_size, num_teacher_layers=config.num_teacher_layers, num_heads=config.num_attention_heads, ffn_size=config.ffn_size, layer_embedding_size=config.layer_embedding_size, dropout=config.dropout, layer_norm_eps=config.layer_norm_eps, ) self.post_init() def forward(self, hidden_states, layer_ids=None, start_layer=None, end_layer=None): if layer_ids is not None: return self.reconstructor(hidden_states, layer_ids) if start_layer is None: start_layer = 0 return self.reconstructor.rollout(hidden_states, start_layer=start_layer, end_layer=end_layer)