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