Girinath11 commited on
Commit
2caeef3
·
verified ·
1 Parent(s): 2512e7d

Update mixture_of_recursion.py

Browse files
Files changed (1) hide show
  1. mixture_of_recursion.py +131 -132
mixture_of_recursion.py CHANGED
@@ -5,7 +5,7 @@ from transformers import PretrainedConfig, PreTrainedModel
5
  from transformers.modeling_outputs import CausalLMOutputWithPast
6
  import math
7
  class RecursiveLanguageModelConfig(PretrainedConfig):
8
- model_type = "recursive_language_model"
9
  def __init__(
10
  self,
11
  vocab_size=50260,
@@ -33,93 +33,93 @@ class RecursiveLanguageModelConfig(PretrainedConfig):
33
  eos_token_id=eos_token_id,
34
  **kwargs
35
  )
36
- self.vocab_size = vocab_size
37
- self.embedding_dim = embedding_dim
38
- self.num_layers = num_layers
39
- self.num_attention_heads = num_attention_heads
40
- self.max_recursion_steps = max_recursion_steps
41
- self.max_position_embeddings = max_position_embeddings
42
- self.hidden_dropout_prob = hidden_dropout_prob
43
- self.attention_dropout_prob = attention_dropout_prob
44
- self.intermediate_size = intermediate_size
45
- self.layer_norm_eps = layer_norm_eps
46
- self.simple_recursion_steps = simple_recursion_steps
47
- self.medium_recursion_steps = medium_recursion_steps
48
- self.complex_recursion_steps = complex_recursion_steps
49
- self.initializer_range = initializer_range
50
  class RotaryPositionalEmbedding(nn.Module):
51
- def __init__(self, dim, max_seq_len=2048, base=10000):
52
  super().__init__()
53
- inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
54
- self.register_buffer('inv_freq', inv_freq)
55
- def forward(self, seq_len, device):
56
- t = torch.arange(seq_len, device=device).float()
57
- freqs = torch.outer(t, self.inv_freq)
58
- emb = torch.cat([freqs, freqs], dim=-1)
59
- return emb.cos(), emb.sin()
60
  def rotate_half(x):
61
- x1, x2 = x[..., :x.shape[-1]//2], x[..., x.shape[-1]//2:]
62
- return torch.cat([-x2, x1], dim=-1)
63
  def apply_rotary_pos_emb(q, k, cos, sin):
64
- cos = cos.unsqueeze(0).unsqueeze(0)
65
- sin = sin.unsqueeze(0).unsqueeze(0)
66
- return (q * cos) + (rotate_half(q) * sin), (k * cos) + (rotate_half(k) * sin)
67
  class MultiHeadAttention(nn.Module):
68
  def __init__(self, config):
69
  super().__init__()
70
- self.num_heads = config.num_attention_heads
71
- self.head_dim = config.embedding_dim // config.num_attention_heads
72
- self.embed_dim = config.embedding_dim
73
- assert self.embed_dim % self.num_heads == 0
74
- self.q_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
75
- self.k_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
76
- self.v_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
77
- self.out_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=False)
78
- self.attn_drop = nn.Dropout(config.attention_dropout_prob)
79
- self.rope = RotaryPositionalEmbedding(self.head_dim, config.max_position_embeddings)
80
- def forward(self, x, causal_mask=None):
81
- B, T, C = x.shape
82
- q = self.q_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
83
- k = self.k_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
84
- v = self.v_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
85
- cos, sin = self.rope(T, x.device)
86
- q, k = apply_rotary_pos_emb(q, k, cos, sin)
87
- scale = math.sqrt(self.head_dim)
88
- scores = torch.matmul(q, k.transpose(-2, -1)) / scale
89
  if causal_mask is not None:
90
- scores = scores + causal_mask
91
- scores = scores.clamp(min=-1e4, max=1e4)
92
- attn = F.softmax(scores, dim=-1)
93
- attn = torch.nan_to_num(attn, nan=0.0, posinf=0.0, neginf=0.0)
94
- attn = self.attn_drop(attn)
95
- out = torch.matmul(attn, v)
96
- out = out.transpose(1, 2).contiguous().view(B, T, C)
97
  return self.out_proj(out)
98
  class FeedForward(nn.Module):
99
  def __init__(self, config):
100
  super().__init__()
101
- self.fc1 = nn.Linear(config.embedding_dim, config.intermediate_size, bias=False)
102
- self.fc2 = nn.Linear(config.intermediate_size, config.embedding_dim, bias=False)
103
- self.drop = nn.Dropout(config.hidden_dropout_prob)
104
  def forward(self, x):
105
  return self.drop(self.fc2(self.drop(F.gelu(self.fc1(x)))))
106
  class TransformerBlock(nn.Module):
107
  def __init__(self, config):
108
  super().__init__()
109
- self.attn = MultiHeadAttention(config)
110
- self.ff = FeedForward(config)
111
- self.ln1 = nn.LayerNorm(config.embedding_dim, eps=config.layer_norm_eps)
112
- self.ln2 = nn.LayerNorm(config.embedding_dim, eps=config.layer_norm_eps)
113
  def forward(self, x, mask=None):
114
- x = x + self.attn(self.ln1(x), mask)
115
- x = x + self.ff(self.ln2(x))
116
  return x
117
  class SequenceLevelRouter(nn.Module):
118
  def __init__(self, config):
119
  super().__init__()
120
- self.pooler = nn.Linear(config.embedding_dim, config.embedding_dim)
121
- self.act = nn.Tanh()
122
- self.head = nn.Sequential(
123
  nn.Linear(config.embedding_dim, config.embedding_dim // 2),
124
  nn.GELU(),
125
  nn.Dropout(0.1),
@@ -132,32 +132,32 @@ class SequenceLevelRouter(nn.Module):
132
  ], dtype=torch.long))
133
  def forward(self, x, valid_mask=None):
134
  if valid_mask is not None:
135
- m = valid_mask.unsqueeze(-1).float()
136
- pooled = (x * m).sum(1) / m.sum(1).clamp(min=1e-9)
137
  else:
138
- pooled = x.mean(1)
139
- pooled = self.act(self.pooler(pooled))
140
- logits = self.head(pooled)
141
- cls = logits.argmax(dim=-1)
142
  return logits, cls, self.steps_map[cls]
143
  class RecursionLayer(nn.Module):
144
  def __init__(self, config):
145
  super().__init__()
146
- self.block = TransformerBlock(config)
147
- def forward(self, x, mask=None):
148
  return self.block(x, mask)
149
  class RecursiveLanguageModel(PreTrainedModel):
150
  config_class = RecursiveLanguageModelConfig
151
  def __init__(self, config):
152
  super().__init__(config)
153
- self.config = config
154
- self.embed = nn.Embedding(config.vocab_size, config.embedding_dim,
155
  padding_idx=config.pad_token_id)
156
- self.layers = nn.ModuleList([TransformerBlock(config) for _ in range(config.num_layers)])
157
- self.router = SequenceLevelRouter(config)
158
- self.rec_layer = RecursionLayer(config)
159
- self.ln_f = nn.LayerNorm(config.embedding_dim, eps=config.layer_norm_eps)
160
- self.lm_head = nn.Linear(config.embedding_dim, config.vocab_size, bias=False)
161
  self.post_init()
162
  def get_input_embeddings(self): return self.embed
163
  def set_input_embeddings(self, v): self.embed = v
@@ -165,7 +165,7 @@ class RecursiveLanguageModel(PreTrainedModel):
165
  def set_output_embeddings(self, v): self.lm_head = v
166
  def _init_weights(self, module):
167
  if isinstance(module, nn.Linear):
168
- nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
169
  if module.bias is not None:
170
  nn.init.zeros_(module.bias)
171
  elif isinstance(module, nn.Embedding):
@@ -176,78 +176,77 @@ class RecursiveLanguageModel(PreTrainedModel):
176
  nn.init.zeros_(module.bias)
177
  nn.init.ones_(module.weight)
178
  def _make_causal_mask(self, input_ids):
179
- B, T = input_ids.shape
180
- device = input_ids.device
181
- mask = torch.zeros(T, T, device=device)
182
- mask = mask.masked_fill(
183
  torch.triu(torch.ones(T, T, device=device, dtype=torch.bool), diagonal=1),
184
  -1e4
185
  )
186
- mask = mask.unsqueeze(0).unsqueeze(0)
187
- pad_mask = (input_ids == self.config.pad_token_id) # [B, T]
188
- valid_mask = ~pad_mask
189
  if pad_mask.any():
190
- pad_key_mask = pad_mask.unsqueeze(1).unsqueeze(2).float() * -1e4
191
- mask = mask + pad_key_mask
192
  return mask, valid_mask
193
- def forward(self, input_ids, labels=None, attention_mask=None, **kwargs):
194
- B, T = input_ids.shape
195
- x = self.embed(input_ids)
196
- causal_mask, valid_mask = self._make_causal_mask(input_ids)
197
  for layer in self.layers:
198
- x = layer(x, causal_mask)
199
- router_logits, cls, steps = self.router(x, valid_mask)
200
- max_steps = int(steps.max().item())
201
  for s in range(max_steps):
202
- gate = (steps > s).float().view(B, 1, 1)
203
- x = gate * self.rec_layer(x, causal_mask) + (1 - gate) * x
204
- x = self.ln_f(x)
205
- logits = self.lm_head(x)
206
- loss = None
207
  if labels is not None:
208
- shift_logits = logits[:, :-1, :].contiguous()
209
- shift_labels = labels[:, 1:].contiguous()
210
- lm_loss = F.cross_entropy(
211
  shift_logits.view(-1, self.config.vocab_size),
212
  shift_labels.view(-1),
213
  ignore_index=-100
214
  )
215
  with torch.no_grad():
216
- per_tok = F.cross_entropy(
217
  shift_logits.view(-1, self.config.vocab_size),
218
  shift_labels.view(-1),
219
  ignore_index=-100, reduction='none'
220
  ).view(B, -1)
221
- valid_tok = (shift_labels != -100).sum(1).clamp(min=1).float()
222
- ppl = torch.exp((per_tok.sum(1) / valid_tok).clamp(max=20))
223
- pseudo = torch.zeros(B, dtype=torch.long, device=input_ids.device)
224
- pseudo[(ppl >= 20) & (ppl < 50)] = 1
225
- pseudo[ppl >= 50] = 2
226
- router_loss = F.cross_entropy(router_logits, pseudo)
227
- loss = lm_loss + 0.1 * router_loss
228
- return CausalLMOutputWithPast(loss=loss, logits=logits)
229
  @torch.no_grad()
230
- def generate(self, input_ids, max_new_tokens=100, temperature=0.8,
231
  top_p=0.9, do_sample=True, **kwargs):
232
  self.eval()
233
- gen = input_ids
234
  for _ in range(max_new_tokens):
235
- ctx = gen[:, -self.config.max_position_embeddings:]
236
- logits = self.forward(ctx).logits[:, -1, :]
237
- if temperature != 1.0:
238
- logits = logits / temperature
239
  if do_sample:
240
- probs = F.softmax(logits, dim=-1)
241
- sorted_probs, sorted_idx = torch.sort(probs, descending=True)
242
- cum_probs = torch.cumsum(sorted_probs, dim=-1)
243
- remove = cum_probs - sorted_probs > top_p
244
- sorted_probs[remove] = 0.0
245
- sorted_probs = sorted_probs / sorted_probs.sum(dim=-1, keepdim=True)
246
- next_tok = torch.gather(sorted_idx, -1,
247
- torch.multinomial(sorted_probs, 1))
248
  else:
249
- next_tok = logits.argmax(dim=-1, keepdim=True)
250
- gen = torch.cat([gen, next_tok], dim=-1)
251
- if (next_tok == self.config.eos_token_id).all():
252
  break
253
  return gen
 
5
  from transformers.modeling_outputs import CausalLMOutputWithPast
6
  import math
7
  class RecursiveLanguageModelConfig(PretrainedConfig):
8
+ model_type="recursive_language_model"
9
  def __init__(
10
  self,
11
  vocab_size=50260,
 
33
  eos_token_id=eos_token_id,
34
  **kwargs
35
  )
36
+ self.vocab_size=vocab_size
37
+ self.embedding_dim=embedding_dim
38
+ self.num_layers=num_layers
39
+ self.num_attention_heads=num_attention_heads
40
+ self.max_recursion_steps=max_recursion_steps
41
+ self.max_position_embeddings=max_position_embeddings
42
+ self.hidden_dropout_prob=hidden_dropout_prob
43
+ self.attention_dropout_prob=attention_dropout_prob
44
+ self.intermediate_size=intermediate_size
45
+ self.layer_norm_eps=layer_norm_eps
46
+ self.simple_recursion_steps=simple_recursion_steps
47
+ self.medium_recursion_steps=medium_recursion_steps
48
+ self.complex_recursion_steps=complex_recursion_steps
49
+ self.initializer_range=initializer_range
50
  class RotaryPositionalEmbedding(nn.Module):
51
+ def __init__(self,dim,max_seq_len=2048,base=10000):
52
  super().__init__()
53
+ inv_freq=1.0/(base**(torch.arange(0,dim,2).float()/dim))
54
+ self.register_buffer('inv_freq',inv_freq)
55
+ def forward(self,seq_len,device):
56
+ t=torch.arange(seq_len,device=device).float()
57
+ freqs=torch.outer(t,self.inv_freq)
58
+ emb=torch.cat([freqs,freqs], dim=-1)
59
+ return emb.cos(),emb.sin()
60
  def rotate_half(x):
61
+ x1, x2=x[..., :x.shape[-1]//2], x[..., x.shape[-1]//2:]
62
+ return torch.cat([-x2, x1],dim=-1)
63
  def apply_rotary_pos_emb(q, k, cos, sin):
64
+ cos=cos.unsqueeze(0).unsqueeze(0)
65
+ sin=sin.unsqueeze(0).unsqueeze(0)
66
+ return(q*cos)+(rotate_half(q)*sin),(k*cos)+(rotate_half(k)*sin)
67
  class MultiHeadAttention(nn.Module):
68
  def __init__(self, config):
69
  super().__init__()
70
+ self.num_heads=config.num_attention_heads
71
+ self.head_dim=config.embedding_dim // config.num_attention_heads
72
+ self.embed_dim=config.embedding_dim
73
+ assert self.embed_dim % self.num_heads==0
74
+ self.q_proj=nn.Linear(self.embed_dim, self.embed_dim, bias=False)
75
+ self.k_proj=nn.Linear(self.embed_dim, self.embed_dim, bias=False)
76
+ self.v_proj=nn.Linear(self.embed_dim, self.embed_dim, bias=False)
77
+ self.out_proj=nn.Linear(self.embed_dim, self.embed_dim, bias=False)
78
+ self.attn_drop=nn.Dropout(config.attention_dropout_prob)
79
+ self.rope=RotaryPositionalEmbedding(self.head_dim, config.max_position_embeddings)
80
+ def forward(self,x,causal_mask=None):
81
+ B,T,C=x.shape
82
+ q=self.q_proj(x).view(B,T,self.num_heads, self.head_dim).transpose(1, 2)
83
+ k=self.k_proj(x).view(B,T,self.num_heads, self.head_dim).transpose(1, 2)
84
+ v=self.v_proj(x).view(B,T,self.num_heads, self.head_dim).transpose(1, 2)
85
+ cos,sin=self.rope(T,x.device)
86
+ q,k=apply_rotary_pos_emb(q, k, cos, sin)
87
+ scale=math.sqrt(self.head_dim)
88
+ scores=torch.matmul(q, k.transpose(-2, -1))/scale
89
  if causal_mask is not None:
90
+ scores=scores+causal_mask
91
+ scores=scores.clamp(min=-1e4, max=1e4)
92
+ attn=F.softmax(scores, dim=-1)
93
+ attn=torch.nan_to_num(attn, nan=0.0, posinf=0.0, neginf=0.0)
94
+ attn=self.attn_drop(attn)
95
+ out=torch.matmul(attn, v)
96
+ out=out.transpose(1, 2).contiguous().view(B, T, C)
97
  return self.out_proj(out)
98
  class FeedForward(nn.Module):
99
  def __init__(self, config):
100
  super().__init__()
101
+ self.fc1=nn.Linear(config.embedding_dim, config.intermediate_size, bias=False)
102
+ self.fc2=nn.Linear(config.intermediate_size, config.embedding_dim, bias=False)
103
+ self.drop=nn.Dropout(config.hidden_dropout_prob)
104
  def forward(self, x):
105
  return self.drop(self.fc2(self.drop(F.gelu(self.fc1(x)))))
106
  class TransformerBlock(nn.Module):
107
  def __init__(self, config):
108
  super().__init__()
109
+ self.attn=MultiHeadAttention(config)
110
+ self.ff=FeedForward(config)
111
+ self.ln1=nn.LayerNorm(config.embedding_dim,eps=config.layer_norm_eps)
112
+ self.ln2=nn.LayerNorm(config.embedding_dim,eps=config.layer_norm_eps)
113
  def forward(self, x, mask=None):
114
+ x=x+self.attn(self.ln1(x), mask)
115
+ x=x+self.ff(self.ln2(x))
116
  return x
117
  class SequenceLevelRouter(nn.Module):
118
  def __init__(self, config):
119
  super().__init__()
120
+ self.pooler=nn.Linear(config.embedding_dim, config.embedding_dim)
121
+ self.act=nn.Tanh()
122
+ self.head=nn.Sequential(
123
  nn.Linear(config.embedding_dim, config.embedding_dim // 2),
124
  nn.GELU(),
125
  nn.Dropout(0.1),
 
132
  ], dtype=torch.long))
133
  def forward(self, x, valid_mask=None):
134
  if valid_mask is not None:
135
+ m=valid_mask.unsqueeze(-1).float()
136
+ pooled=(x * m).sum(1)/m.sum(1).clamp(min=1e-9)
137
  else:
138
+ pooled=x.mean(1)
139
+ pooled=self.act(self.pooler(pooled))
140
+ logits=self.head(pooled)
141
+ cls=logits.argmax(dim=-1)
142
  return logits, cls, self.steps_map[cls]
143
  class RecursionLayer(nn.Module):
144
  def __init__(self, config):
145
  super().__init__()
146
+ self.block=TransformerBlock(config)
147
+ def forward(self,x,mask=None):
148
  return self.block(x, mask)
149
  class RecursiveLanguageModel(PreTrainedModel):
150
  config_class = RecursiveLanguageModelConfig
151
  def __init__(self, config):
152
  super().__init__(config)
153
+ self.config=config
154
+ self.embed=nn.Embedding(config.vocab_size, config.embedding_dim,
155
  padding_idx=config.pad_token_id)
156
+ self.layers=nn.ModuleList([TransformerBlock(config) for _ in range(config.num_layers)])
157
+ self.router=SequenceLevelRouter(config)
158
+ self.rec_layer=RecursionLayer(config)
159
+ self.ln_f=nn.LayerNorm(config.embedding_dim, eps=config.layer_norm_eps)
160
+ self.lm_head=nn.Linear(config.embedding_dim, config.vocab_size, bias=False)
161
  self.post_init()
162
  def get_input_embeddings(self): return self.embed
163
  def set_input_embeddings(self, v): self.embed = v
 
165
  def set_output_embeddings(self, v): self.lm_head = v
166
  def _init_weights(self, module):
167
  if isinstance(module, nn.Linear):
168
+ nn.init.normal_(module.weight,mean=0.0,std=self.config.initializer_range)
169
  if module.bias is not None:
170
  nn.init.zeros_(module.bias)
171
  elif isinstance(module, nn.Embedding):
 
176
  nn.init.zeros_(module.bias)
177
  nn.init.ones_(module.weight)
178
  def _make_causal_mask(self, input_ids):
179
+ B,T=input_ids.shape
180
+ device=input_ids.device
181
+ mask=torch.zeros(T, T, device=device)
182
+ mask=mask.masked_fill(
183
  torch.triu(torch.ones(T, T, device=device, dtype=torch.bool), diagonal=1),
184
  -1e4
185
  )
186
+ mask=mask.unsqueeze(0).unsqueeze(0)
187
+ pad_mask=(input_ids==self.config.pad_token_id)
188
+ valid_mask=~pad_mask
189
  if pad_mask.any():
190
+ pad_key_mask=pad_mask.unsqueeze(1).unsqueeze(2).float()*-1e4
191
+ mask=mask+pad_key_mask
192
  return mask, valid_mask
193
+ def forward(self,input_ids,labels=None,attention_mask=None,**kwargs):
194
+ B,T=input_ids.shape
195
+ x=self.embed(input_ids)
196
+ causal_mask,valid_mask=self._make_causal_mask(input_ids)
197
  for layer in self.layers:
198
+ x=layer(x,causal_mask)
199
+ router_logits, cls, steps=self.router(x,valid_mask)
200
+ max_steps=int(steps.max().item())
201
  for s in range(max_steps):
202
+ gate=(steps > s).float().view(B, 1, 1)
203
+ x=gate*self.rec_layer(x, causal_mask)+(1-gate)*x
204
+ x=self.ln_f(x)
205
+ logits=self.lm_head(x)
206
+ loss=None
207
  if labels is not None:
208
+ shift_logits=logits[:, :-1, :].contiguous()
209
+ shift_labels=labels[:, 1:].contiguous()
210
+ lm_loss=F.cross_entropy(
211
  shift_logits.view(-1, self.config.vocab_size),
212
  shift_labels.view(-1),
213
  ignore_index=-100
214
  )
215
  with torch.no_grad():
216
+ per_tok= F.cross_entropy(
217
  shift_logits.view(-1, self.config.vocab_size),
218
  shift_labels.view(-1),
219
  ignore_index=-100, reduction='none'
220
  ).view(B, -1)
221
+ valid_tok=(shift_labels!=-100).sum(1).clamp(min=1).float()
222
+ ppl=torch.exp((per_tok.sum(1)/valid_tok).clamp(max=20))
223
+ pseudo=torch.zeros(B,dtype=torch.long,device=input_ids.device)
224
+ pseudo[(ppl>=20)&(ppl<50)]=1
225
+ pseudo[ppl>= 50] = 2
226
+ router_loss=F.cross_entropy(router_logits, pseudo)
227
+ loss=lm_loss+0.1*router_loss
228
+ return CausalLMOutputWithPast(loss=loss,logits=logits)
229
  @torch.no_grad()
230
+ def generate(self,input_ids,max_new_tokens=100,temperature=0.8,
231
  top_p=0.9, do_sample=True, **kwargs):
232
  self.eval()
233
+ gen=input_ids
234
  for _ in range(max_new_tokens):
235
+ ctx=gen[:,-self.config.max_position_embeddings:]
236
+ logits=self.forward(ctx).logits[:, -1, :]
237
+ if temperature!=1.0:
238
+ logits=logits/temperature
239
  if do_sample:
240
+ probs=F.softmax(logits, dim=-1)
241
+ sorted_probs,sorted_idx=torch.sort(probs,descending=True)
242
+ cum_probs=torch.cumsum(sorted_probs, dim=-1)
243
+ remove=cum_probs-sorted_probs>top_p
244
+ sorted_probs[remove]=0.0
245
+ sorted_probs=sorted_probs/sorted_probs.sum(dim=-1,keepdim=True)
246
+ next_tok=torch.gather(sorted_idx,-1,torch.multinomial(sorted_probs,1))
 
247
  else:
248
+ next_tok=logits.argmax(dim=-1,keepdim=True)
249
+ gen=torch.cat([gen,next_tok],dim=-1)
250
+ if (next_tok==self.config.eos_token_id).all():
251
  break
252
  return gen