homeGPT / modeling_shared_reconstructor.py
summerMC's picture
Upload folder using huggingface_hub
f4f7ccd verified
Raw
History Blame Contribute Delete
5.12 kB
# 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)