File size: 5,124 Bytes
f4f7ccd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
# 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)