leejunhyeok commited on
Commit
f362538
·
verified ·
1 Parent(s): 7b041df

Update modeling_motif.py

Browse files
Files changed (1) hide show
  1. modeling_motif.py +23 -728
modeling_motif.py CHANGED
@@ -1,5 +1,5 @@
1
  import math
2
- from typing import List, Optional, Tuple, Union
3
 
4
  import torch
5
  import torch.utils.checkpoint
@@ -35,144 +35,22 @@ logger = logging.get_logger(__name__)
35
  if is_flash_attn_2_available():
36
  from transformers.modeling_flash_attention_utils import _flash_attention_forward
37
 
38
- try:
39
- moreh_ops = torch.ops.moreh
40
- MorehRMSNorm = moreh_ops.T5LayerNorm
41
- ScaledDotProductAttention = moreh_ops.scaled_dot_product_attention
42
- MorehFlashAttention = moreh_ops.flash_attention
43
- logger.warning_once("Using moreh ops")
44
- except AttributeError:
45
- MorehRMSNorm = None
46
- ScaledDotProductAttention = None
47
- MorehFlashAttention = None
48
- logger.warning_once("Failed to import moreh ops")
49
-
50
-
51
-
52
- # DEBUG = False
53
- # logger.info(f"DEBUG: {DEBUG} : will log timing")
54
- # def log_timing(obj):
55
- # """Decorator to log timing of function or class execution"""
56
- # if isinstance(obj, type):
57
- # # If decorating a class
58
- # class TimedClass(obj):
59
- # def __getattribute__(self, name):
60
- # attr = super().__getattribute__(name)
61
- # if callable(attr) and not name.startswith('__'):
62
- # def timed_method(*args, **kwargs):
63
- # if not DEBUG:
64
- # return attr(*args, **kwargs)
65
- # if name != "forward":
66
- # return attr(*args, **kwargs)
67
-
68
- # start_time = time.time()
69
- # logger.info(f"Entering {obj.__name__}.{name}")
70
- # result = attr(*args, **kwargs)
71
- # end_time = time.time()
72
- # logger.info(f"Exiting {obj.__name__}.{name}, took {end_time - start_time:.4f} seconds")
73
- # return result
74
- # return timed_method
75
- # return attr
76
- # return TimedClass
77
- # else:
78
- # # If decorating a function
79
- # def wrapper(*args, **kwargs):
80
- # if not DEBUG:
81
- # return obj(*args, **kwargs)
82
-
83
- # start_time = time.time()
84
- # logger.info(f"Entering {obj.__name__}")
85
- # result = obj(*args, **kwargs)
86
- # end_time = time.time()
87
- # logger.info(f"Exiting {obj.__name__}, took {end_time - start_time:.4f} seconds")
88
- # return result
89
- # return wrapper
90
-
91
-
92
 
93
  #_CHECKPOINT_FOR_DOC = "moreh/Motif-102B"
94
  _CONFIG_FOR_DOC = "MotifConfig"
95
 
96
- #from .moreh_moe import MorehMoeMLP, MorehMoeFusedMLP
97
-
98
- import torch
99
  from transformers.activations import ACT2CLS as _ACT2CLS
100
  from transformers.activations import ClassInstantier
101
- moreh_ops = torch.ops.moreh
102
-
103
- from typing import Callable, Dict, List, Tuple
104
-
105
- import torch
106
-
107
-
108
- # @log_timing
109
- def multi_head_forward_backward(shared_activation: torch.Tensor,
110
- head_fns: List[Callable[[torch.Tensor], Dict[str, torch.Tensor]]],
111
- return_keys=("loss", ),
112
- return_only_first_head=True) -> Tuple[torch.Tensor, ...]:
113
- """
114
- The forward-backward pattern introduced in the paper https://arxiv.org/abs/2404.19737
115
- to reduce memory overhead due to activations from multiple heads.
116
-
117
- Args:
118
- - shared_activation: the shared activation across all heads
119
- - head_fns: the head-wise forward computations that start from `shared_activation`.
120
- it should output a dictionary of tensors with keys matching `return_keys`
121
- - return_keys: the keys to return in order
122
- - return_only_first_head: whether to return only the values from the first head
123
-
124
- Returns:
125
- - a tuple of return tensors
126
-
127
- Side effect:
128
- - (only when `torch.is_grad_enabled()`)
129
- the gradients accumulated as if `sum(head_fn(shared_activation)["loss"] for head_fn in head_fns).backward()` had been called
130
- """
131
- if not return_only_first_head:
132
- raise NotImplementedError
133
-
134
- return_key_set = set(return_keys)
135
- if "loss" not in return_key_set:
136
- raise Exception("'loss' is a required return key.")
137
-
138
- detached_shared_activation = shared_activation.detach()
139
- detached_shared_activation.requires_grad = True
140
- return_values = {key: None for key in return_keys}
141
- for head_idx, head_fn in enumerate(head_fns):
142
- if head_idx > 0 and not torch.is_grad_enabled():
143
- continue
144
-
145
- # forward pass for the head
146
- headwise_outputs = head_fn(detached_shared_activation)
147
- if set(headwise_outputs.keys()) != return_key_set:
148
- raise Exception(f"Headwise output keys {headwise_outputs.keys()} do not match return keys {return_keys}.")
149
-
150
- # backward pass for the head
151
- # effect 1: the parameters of the head
152
- # effect 2: gradient accumulated in `detached_shared_activation.grad`
153
- if torch.is_grad_enabled():
154
- headwise_loss = headwise_outputs["loss"]
155
- headwise_loss.backward(
156
- ) # NOTE: You do not need to retain graph since no graph is shared across backward passes
157
-
158
- if head_idx == 0:
159
- for key in return_keys:
160
- return_values[key] = headwise_outputs[key]
161
-
162
- assert all(value is not None for value in return_values.values())
163
-
164
- # backward pass for the shared part
165
- if torch.is_grad_enabled():
166
- shared_activation.backward(detached_shared_activation.grad)
167
-
168
- return tuple(return_values[key] for key in return_keys)
169
 
170
 
171
  class PolyNorm(torch.nn.Module):
172
  """
173
  A trainable activation function introduced in https://arxiv.org/html/2411.03884v1.
174
  The code is copied from https://github.com/BryceZhuo/PolyCom?tab=readme-ov-file/README.md,
175
- with the change `* torch.rsqrt` => `/ torch.sqrt` for potential MAF incompatibility.
176
  """
177
 
178
  def __init__(self, eps=1e-6):
@@ -189,31 +67,11 @@ class PolyNorm(torch.nn.Module):
189
  x ** 2) + self.weight[2] * self._norm(x) + self.bias
190
 
191
 
192
- class PolyNorm_Test(torch.nn.Module):
193
- """
194
- A trainable activation function introduced in https://arxiv.org/html/2411.03884v1.
195
- The code is copied from https://github.com/BryceZhuo/PolyCom?tab=readme-ov-file/README.md,
196
- with the change `* torch.rsqrt` => `/ torch.sqrt` for potential MAF incompatibility.
197
- """
198
-
199
- def __init__(self, eps=1e-6):
200
- super(PolyNorm_Test, self).__init__()
201
- self.weight = torch.nn.Parameter(torch.ones(3) / 3)
202
- self.bias = torch.nn.Parameter(torch.zeros(1))
203
- self.eps = eps
204
-
205
- def forward(self, x):
206
-
207
- #return torch.nn.SiLU(x)
208
- return moreh_ops.poly_norm(x, self.weight, self.bias)
209
-
210
-
211
- CUSTOM_ACT2CLS = {"poly_norm": PolyNorm_Test, "poly_norm_test": PolyNorm_Test}
212
  ACT2CLS = {**_ACT2CLS, **CUSTOM_ACT2CLS}
213
  ACT2FN = ClassInstantier(ACT2CLS)
214
 
215
 
216
-
217
  class MotifRMSNorm(nn.Module):
218
 
219
  def __init__(self, hidden_size, eps=1e-6):
@@ -226,8 +84,7 @@ class MotifRMSNorm(nn.Module):
226
 
227
  def forward(self, hidden_states):
228
  input_dtype = hidden_states.dtype
229
- hidden_states = hidden_states.to(torch.float32)
230
- variance = hidden_states.pow(2).mean(-1, keepdim=True)
231
  hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
232
  return self.weight * hidden_states.to(input_dtype)
233
 
@@ -235,7 +92,7 @@ class MotifRMSNorm(nn.Module):
235
  return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"
236
 
237
 
238
- ALL_LAYERNORM_LAYERS.append(MotifRMSNorm if MorehRMSNorm is None else MorehRMSNorm)
239
 
240
 
241
  class MotifRotaryEmbeddingWithCache(nn.Module):
@@ -293,7 +150,6 @@ class MotifRotaryEmbeddingWithCache(nn.Module):
293
  )
294
 
295
 
296
- # @log_timing
297
  class MotifRotaryEmbedding(nn.Module):
298
 
299
  def __init__(
@@ -307,7 +163,6 @@ class MotifRotaryEmbedding(nn.Module):
307
  config: Optional[MotifConfig] = None,
308
  ):
309
  super().__init__()
310
- # TODO (joao): remove the `if` below, only used for BC
311
  self.rope_kwargs = {}
312
  if config is None:
313
  logger.warning_once(
@@ -352,7 +207,7 @@ class MotifRotaryEmbedding(nn.Module):
352
  device,
353
  seq_len=seq_len,
354
  **self.rope_kwargs)
355
- self.register_buffer("inv_freq", inv_freq, persistent=False) # TODO joao: may break with compilation
356
  self.max_seq_len_cached = seq_len
357
 
358
  if seq_len < self.original_max_seq_len and self.max_seq_len_cached > self.original_max_seq_len: # reset
@@ -401,7 +256,6 @@ def rotate_half(x):
401
  return rotated_tensor
402
 
403
 
404
- # @log_timing
405
  def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1, fused_rope=True):
406
  """
407
  Applies rotary position embeddings to the input tensors.
@@ -427,12 +281,7 @@ def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1, fus
427
  q_embed = (q * cos) + (rotate_half(q) * sin)
428
  k_embed = (k * cos) + (rotate_half(k) * sin)
429
  '''
430
- #cos = cos[position_ids]
431
- #sin = sin[position_ids]
432
-
433
- #cos = cos[position_ids].unsqueeze(unsqueeze_dim) # [bs, 1, seq_len, dim]
434
- #sin = sin[position_ids].unsqueeze(unsqueeze_dim) # [bs, 1, seq_len, dim]
435
-
436
  q = q.transpose(1, 2)
437
  k = k.transpose(1, 2)
438
 
@@ -450,9 +299,8 @@ def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1, fus
450
  return q_embed, k_embed
451
 
452
 
453
- # @log_timing
454
  class MotifMLP(nn.Module):
455
-
456
  def __init__(self, config):
457
  super().__init__()
458
  self.hidden_size = config.hidden_size
@@ -462,394 +310,11 @@ class MotifMLP(nn.Module):
462
  self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
463
  self.act_fn = ACT2FN[config.hidden_act]
464
 
465
- if config.wesar_weights:
466
- self.gate_up_proj_alpha = nn.Parameter(torch.tensor(1) *config.gate_up_proj_alpha)
467
- self.down_proj_alpha = nn.Parameter(torch.tensor(1) * config.down_proj_alpha)
468
- else:
469
- self.gate_up_proj_alpha=1
470
- self.down_proj_alpha=1
471
- if config.muP:
472
- self.down_proj.__do_scale_tager__ = True
473
- self.gate_proj.__do_scale_tager_mu_dim_model__ = True
474
- self.up_proj.__do_scale_tager_mu_dim_model__ = True
475
- self.down_proj.__do_scale_tager_mu_ffn__ = True
476
-
477
-
478
  def forward(self, hidden_state):
479
- hidden_state = hidden_state*self.gate_up_proj_alpha
480
- #hidden_state = self.down_proj(self.act_fn(self.gate_proj(hidden_state)) * self.up_proj(hidden_state))*
481
- return self.down_proj_alpha*self.down_proj(self.act_fn(self.gate_proj(hidden_state)) * self.up_proj(hidden_state))
482
-
483
-
484
- class MorehMoeFusedMLP(nn.Module):
485
- def __init__(self,
486
- ffn_dim,
487
- hidden_dim,
488
- hidden_act_moe,
489
- num_experts,
490
- num_groups=1,
491
- device=None,
492
- continual_training=False):
493
- super().__init__()
494
- self.ffn_dim = ffn_dim
495
- self.hidden_dim = hidden_dim
496
- self.hidden_act_moe = hidden_act_moe
497
-
498
- self.num_experts = num_experts
499
- self.num_groups = num_groups
500
-
501
- assert self.num_experts % self.num_groups == 0
502
- self.num_experts_per_group = self.num_experts // self.num_groups
503
-
504
- ## bsz, seq, group size, 2*ffn_size
505
-
506
- moreh_ops = torch.ops.moreh
507
- self.w13 = nn.ModuleList([
508
- moreh_ops.MoeFanInLinear(self.hidden_dim,
509
- self.ffn_dim * 2,
510
- bias=False,
511
- num_experts=self.num_experts_per_group,
512
- device=device)
513
- for _ in range(self.num_groups)
514
- ])
515
-
516
- self.w2 = nn.ModuleList([
517
- moreh_ops.MoeFanOutLinear(self.ffn_dim,
518
- self.hidden_dim,
519
- bias=False,
520
- num_experts=self.num_experts_per_group,
521
- device=device)
522
- for _ in range(self.num_groups)
523
- ])
524
-
525
- ## use silu?
526
- self.act_fn = ACT2FN[self.hidden_act_moe]
527
-
528
- if continual_training:
529
- logger.info('two optipons 1. zero init all weights, 2. add scaling param to moe output.')
530
- self._zero_init()
531
-
532
- def _zero_init(self):
533
- for module in self.w2:
534
- for n,param in module.named_parameters():
535
- logger.info(f'{n} {param.shape}')
536
- param.data.zero_()
537
-
538
-
539
- def forward(self, hidden_states, selected_experts, routing_weights):
540
- w13_final_output = None
541
- for group_idx in range(self.num_groups):
542
- w13_output_in_group = self._get_w13_output(hidden_states,
543
- selected_experts,
544
- group_idx)
545
- if w13_final_output is None:
546
- w13_final_output = w13_output_in_group
547
- else:
548
- w13_final_output += w13_output_in_group
549
-
550
- current_hidden_states = self.act_fn(
551
- w13_final_output[:, :, :, :self.ffn_dim]
552
- ) * w13_final_output[:, :, :, self.ffn_dim:]
553
-
554
- final_hidden_states = None
555
- for group_idx in range(self.num_groups):
556
- w2_output_in_group = self._get_w2_output(current_hidden_states,
557
- selected_experts,
558
- routing_weights, group_idx)
559
- if final_hidden_states is None:
560
- final_hidden_states = w2_output_in_group
561
- else:
562
- final_hidden_states += w2_output_in_group
563
- return final_hidden_states
564
-
565
- def _get_w13_output(self, hidden_states, selected_experts, group_idx):
566
- selected_experts_in_group = selected_experts - (
567
- group_idx * self.num_experts_per_group)
568
-
569
- w13_output = self.w13[group_idx](hidden_states,
570
- selected_experts_in_group)
571
- return w13_output
572
-
573
- def _get_w2_output(self, hidden_states, selected_experts, routing_weights,
574
- group_idx):
575
- selected_experts_in_group = selected_experts - (
576
- group_idx * self.num_experts_per_group)
577
- output = self.w2[group_idx](hidden_states, selected_experts_in_group,
578
- routing_weights)
579
- return output
580
-
581
-
582
- class MoEGate(nn.Module):
583
-
584
- def __init__(self, config):
585
- super().__init__()
586
- self.config = config
587
- self.top_k = config.num_experts_per_tok
588
- self.n_routed_experts = config.n_routed_experts
589
- self.routed_scaling_factor = config.routed_scaling_factor
590
- self.scoring_func = config.scoring_func
591
- self.seq_aux = config.seq_aux
592
- self.topk_method = config.topk_method
593
- self.n_group = config.n_group
594
- self.topk_group = config.topk_group
595
-
596
- # topk selection algorithm
597
- self.norm_topk_prob = config.norm_topk_prob
598
- self.gating_dim = config.hidden_size
599
- self.weight = nn.Parameter(
600
- torch.empty((self.n_routed_experts, self.gating_dim)))
601
- if self.topk_method == "noaux_tc":
602
- self.e_score_correction_bias = nn.Parameter(
603
- torch.empty((self.n_routed_experts)))
604
- self.reset_parameters()
605
-
606
- def reset_parameters(self) -> None:
607
- import torch.nn.init as init
608
-
609
- init.kaiming_uniform_(self.weight, a=math.sqrt(5))
610
-
611
- def forward(self, hidden_states):
612
- bsz, seq_len, h = hidden_states.shape
613
- ### compute gating score
614
- hidden_states = hidden_states.view(-1, h)
615
- logits = F.linear(hidden_states.type(torch.float32),
616
- self.weight.type(torch.float32), None)
617
- if self.scoring_func == "sigmoid":
618
- scores = logits.sigmoid()
619
- else:
620
- raise NotImplementedError(
621
- f"insupportable scoring function for MoE gating: {self.scoring_func}"
622
- )
623
-
624
- ### select top-k experts
625
- if self.topk_method == "greedy":
626
- topk_weight, topk_idx = torch.topk(scores,
627
- k=self.top_k,
628
- dim=-1,
629
- sorted=False)
630
- elif self.topk_method == "group_limited_greedy":
631
- group_scores = (scores.view(bsz * seq_len, self.n_group,
632
- -1).max(dim=-1).values) # [n, n_group]
633
- group_idx = torch.topk(group_scores,
634
- k=self.topk_group,
635
- dim=-1,
636
- sorted=False)[1] # [n, top_k_group]
637
- group_mask = torch.zeros_like(group_scores) # [n, n_group]
638
- group_mask.scatter_(1, group_idx, 1) # [n, n_group]
639
- score_mask = (group_mask.unsqueeze(-1).expand(
640
- bsz * seq_len, self.n_group,
641
- self.n_routed_experts // self.n_group).reshape(
642
- bsz * seq_len, -1)) # [n, e]
643
- tmp_scores = scores.masked_fill(~score_mask.bool(), 0.0) # [n, e]
644
- topk_weight, topk_idx = torch.topk(tmp_scores,
645
- k=self.top_k,
646
- dim=-1,
647
- sorted=False)
648
- elif self.topk_method == "noaux_tc":
649
- ### will be used. ###
650
- scores_for_choice = scores.view(
651
- bsz * seq_len, -1) + self.e_score_correction_bias.unsqueeze(0)
652
- group_scores = (scores_for_choice.view(
653
- bsz * seq_len, self.n_group,
654
- -1).topk(2, dim=-1)[0].sum(dim=-1)) # [n, n_group]
655
- group_idx = torch.topk(group_scores,
656
- k=self.topk_group,
657
- dim=-1,
658
- sorted=False)[1] # [n, top_k_group]
659
- group_mask = torch.zeros_like(group_scores) # [n, n_group]
660
- group_mask.scatter_(1, group_idx, 1) # [n, n_group]
661
- score_mask = (group_mask.unsqueeze(-1).expand(
662
- bsz * seq_len, self.n_group,
663
- self.n_routed_experts // self.n_group).reshape(
664
- bsz * seq_len, -1)) # [n, e]
665
- tmp_scores = scores_for_choice.masked_fill(~score_mask.bool(),
666
- 0.0) # [n, e]
667
- _, topk_idx = torch.topk(tmp_scores,
668
- k=self.top_k,
669
- dim=-1,
670
- sorted=False)
671
- topk_weight = scores.gather(1, topk_idx)
672
- else:
673
- raise NotImplementedError(
674
- f"insupportable TopK function for MoE gating: {self.topk_method}"
675
- )
676
-
677
- ### norm gate to sum 1
678
- if self.top_k > 1 and self.norm_topk_prob:
679
- denominator = topk_weight.sum(dim=-1, keepdim=True) + 1e-20
680
- topk_weight = topk_weight / denominator
681
- topk_weight = topk_weight * self.routed_scaling_factor # must multiply the scaling factor
682
-
683
- return topk_idx, topk_weight
684
-
685
-
686
- class MotifMoE(nn.Module):
687
- """
688
- A mixed expert module containing shared experts.
689
- """
690
- def __init__(self, config):
691
- super().__init__()
692
- self.config = config
693
- self.num_experts_per_tok = config.num_experts_per_tok
694
- self.use_moreh_moe = config.use_moreh_moe
695
- self.use_fused_mlp = config.use_fused_mlp
696
-
697
- if hasattr(config, "ep_size") and config.ep_size > 1:
698
- assert config.ep_size == dist.get_world_size()
699
- assert not config.use_moreh_moe
700
- self.ep_size = config.ep_size
701
- self.experts_per_rank = config.n_routed_experts // config.ep_size
702
- self.ep_rank = dist.get_rank()
703
- self.experts = nn.ModuleList([
704
- (DeepseekV3MLP(config,
705
- intermediate_size=config.moe_intermediate_size)
706
- if i >= self.ep_rank * self.experts_per_rank and i <
707
- (self.ep_rank + 1) * self.experts_per_rank else None)
708
- for i in range(config.n_routed_experts)
709
- ])
710
- else:
711
- self.ep_size = 1
712
- self.experts_per_rank = config.n_routed_experts
713
- self.ep_rank = 0
714
- if self.use_moreh_moe:
715
- if not self.use_fused_mlp:
716
- self.experts = MorehMoeMLP(
717
- ffn_dim=config.moe_intermediate_size,
718
- hidden_dim=config.hidden_size,
719
- hidden_act_moe=config.hidden_act_moe,
720
- num_experts=config.n_routed_experts,
721
- device=None)
722
- else:
723
- ## group expert.
724
- self.experts = MorehMoeFusedMLP(
725
- ffn_dim=config.moe_intermediate_size,
726
- hidden_dim=config.hidden_size,
727
- hidden_act_moe=config.hidden_act_moe,
728
- num_experts=config.n_routed_experts,
729
- num_groups=config.n_group,
730
- device=None,
731
- continual_training=config.continual_training,
732
- )
733
- else:
734
- self.experts = nn.ModuleList([
735
- DeepseekV3MLP(
736
- config, intermediate_size=config.moe_intermediate_size)
737
- for i in range(config.n_routed_experts)
738
- ])
739
-
740
- self.gate = MoEGate(config)
741
-
742
- def forward(self, hidden_states):
743
- identity = hidden_states
744
- orig_shape = hidden_states.shape
745
- topk_idx, topk_weight = self.gate(hidden_states)
746
- if self.use_moreh_moe:
747
- y = self.experts(hidden_states, topk_idx.view(*orig_shape[:-1], -1),
748
- topk_weight.view(*orig_shape[:-1], -1))
749
- y = y.type(hidden_states.dtype)
750
- else:
751
- hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
752
- flat_topk_idx = topk_idx.view(-1)
753
- if self.training:
754
- hidden_states = hidden_states.repeat_interleave(
755
- self.num_experts_per_tok, dim=0)
756
- y = torch.empty_like(hidden_states)
757
- for i, expert in enumerate(self.experts):
758
- y[flat_topk_idx == i] = expert(
759
- hidden_states[flat_topk_idx == i])
760
- y = (y.view(*topk_weight.shape, -1) *
761
- topk_weight.unsqueeze(-1)).sum(dim=1)
762
- y = y.type(hidden_states.dtype)
763
- y = y.view(*orig_shape)
764
- # y = AddAuxiliaryLoss.apply(y, aux_loss)
765
- else:
766
- y = self.moe_infer(hidden_states, topk_idx,
767
- topk_weight).view(*orig_shape)
768
- return y, identity
769
-
770
- @torch.no_grad()
771
- def moe_infer(self, x, topk_ids, topk_weight):
772
- cnts = topk_ids.new_zeros((topk_ids.shape[0], len(self.experts)))
773
- cnts.scatter_(1, topk_ids, 1)
774
- tokens_per_expert = cnts.sum(dim=0)
775
- idxs = topk_ids.view(-1).argsort()
776
- sorted_tokens = x[idxs // topk_ids.shape[1]]
777
- sorted_tokens_shape = sorted_tokens.shape
778
- if self.ep_size > 1:
779
- tokens_per_ep_rank = tokens_per_expert.view(self.ep_size,
780
- -1).sum(dim=1)
781
- tokens_per_expert_group = tokens_per_expert.new_empty(
782
- tokens_per_expert.shape[0])
783
- dist.all_to_all_single(tokens_per_expert_group, tokens_per_expert)
784
- output_splits = (tokens_per_expert_group.view(
785
- self.ep_size, -1).sum(1).cpu().numpy().tolist())
786
- gathered_tokens = sorted_tokens.new_empty(
787
- tokens_per_expert_group.sum(dim=0).cpu().item(),
788
- sorted_tokens.shape[1])
789
- input_split_sizes = tokens_per_ep_rank.cpu().numpy().tolist()
790
- dist.all_to_all(
791
- list(gathered_tokens.split(output_splits)),
792
- list(sorted_tokens.split(input_split_sizes)),
793
- )
794
- tokens_per_expert_post_gather = tokens_per_expert_group.view(
795
- self.ep_size, self.experts_per_rank).sum(dim=0)
796
- gatherd_idxs = np.zeros(shape=(gathered_tokens.shape[0],),
797
- dtype=np.int32)
798
- s = 0
799
- for i, k in enumerate(tokens_per_expert_group.cpu().numpy()):
800
- gatherd_idxs[s:s + k] = i % self.experts_per_rank
801
- s += k
802
- gatherd_idxs = gatherd_idxs.argsort()
803
- sorted_tokens = gathered_tokens[gatherd_idxs]
804
- tokens_per_expert = tokens_per_expert_post_gather
805
- tokens_per_expert = tokens_per_expert.cpu().numpy()
806
-
807
- outputs = []
808
- start_idx = 0
809
- for i, num_tokens in enumerate(tokens_per_expert):
810
- end_idx = start_idx + num_tokens
811
- if num_tokens == 0:
812
- continue
813
- expert = self.experts[i + self.ep_rank * self.experts_per_rank]
814
- tokens_for_this_expert = sorted_tokens[start_idx:end_idx]
815
- expert_out = expert(tokens_for_this_expert)
816
- outputs.append(expert_out)
817
- start_idx = end_idx
818
-
819
- outs = torch.cat(outputs,
820
- dim=0) if len(outputs) else sorted_tokens.new_empty(0)
821
- if self.ep_size > 1:
822
- new_x = torch.empty_like(outs)
823
- new_x[gatherd_idxs] = outs
824
- gathered_tokens = new_x.new_empty(*sorted_tokens_shape)
825
- dist.all_to_all(
826
- list(gathered_tokens.split(input_split_sizes)),
827
- list(new_x.split(output_splits)),
828
- )
829
- outs = gathered_tokens
830
-
831
- new_x = torch.empty_like(outs)
832
- new_x[idxs] = outs
833
- final_out = (new_x.view(
834
- *topk_ids.shape, -1).type(topk_weight.dtype).mul_(
835
- topk_weight.unsqueeze(dim=-1)).sum(dim=1).type(new_x.dtype))
836
- return final_out
837
 
838
 
839
  def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
840
-
841
-
842
- """
843
- This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
844
- num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
845
-
846
- batch, num_key_value_heads, slen, head_dim = hidden_states.shape
847
- if n_rep == 1:
848
- return hidden_states
849
- hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
850
- return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
851
- """
852
-
853
  return torch.repeat_interleave(hidden_states, dim=1, repeats=n_rep)
854
 
855
 
@@ -1384,7 +849,6 @@ MOTIF_ATTENTION_CLASSES = {
1384
  }
1385
 
1386
 
1387
- # @log_timing
1388
  class MotifDecoderLayer(nn.Module):
1389
 
1390
  def __init__(self, config: MotifConfig, moe_layer: bool, layer_idx: int):
@@ -1655,7 +1119,6 @@ MOTIF_INPUTS_DOCSTRING = r"""
1655
  """
1656
 
1657
 
1658
- # @log_timing
1659
  @add_start_docstrings(
1660
  "The bare Motif Model outputting raw hidden-states without any specific head on top.",
1661
  MOTIF_START_DOCSTRING,
@@ -1675,17 +1138,13 @@ class MotifModel(MotifPreTrainedModel):
1675
  self.multi_token_heads = config.multi_token_heads
1676
 
1677
  self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
1678
- # NOTE: For multi-token models, the last decoder layers (one for each token index)
1679
- # are implemented as a part of `MotifModelForCausalLM` to enable a custom forward-backward procedure.
1680
 
1681
  num_hidden_layers = config.num_hidden_layers if self.multi_token_heads is None else config.num_hidden_layers - 1
1682
- if config.moe:
1683
- moe_layer = [True for i in range(num_hidden_layers)]
1684
- else:
1685
- moe_layer = [False for i in range(num_hidden_layers)]
1686
  logger.info(f'current_moe layer { moe_layer }')
1687
- self.layers = nn.ModuleList([MotifDecoderLayer(config = config, moe_layer= moe_layer[layer_idx],
1688
- layer_idx=layer_idx) for layer_idx in range(num_hidden_layers)])
 
1689
  self._attn_implementation = config._attn_implementation
1690
  RMSNorm = MorehRMSNorm if MorehRMSNorm is not None else MotifRMSNorm
1691
  self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
@@ -1701,36 +1160,6 @@ class MotifModel(MotifPreTrainedModel):
1701
  self.gradient_checkpointing = False
1702
  self.post_init()
1703
 
1704
- self.use_pipeline = config.use_pipeline
1705
- if self.use_pipeline:
1706
- logger.info('use reinforced pp..')
1707
- if config.num_stages==2:
1708
- ### moe version
1709
- if config.decontam_attn:
1710
- self.split_layers = [15]
1711
- else:
1712
- if num_hidden_layers == 32:
1713
- self.split_layers = [15] # 14: 15,17 # 13: 14:18
1714
- else:
1715
- self.split_layers = [6]
1716
- elif config.num_stages==3:
1717
- self.split_layers = [9,20] ## 11, 11, 10
1718
- elif config.num_stages==4:
1719
- self.split_layers = [7,15,23] #7,9,9,7
1720
- elif config.num_stages==16:
1721
- self.split_layers = [1,3,5,7,9,11,13,15,17,19,21,23,25,27,29]
1722
- logger.info(f' check the split layers (moe): {self.split_layers}')
1723
-
1724
- self.scale_emb = 1
1725
-
1726
- # Reparameterization <|_1_|>
1727
- if config.wesar_weights :
1728
- logger.info(f'config.wesar_weights {config.wesar_weights}')
1729
- self.norm_alpha = nn.Parameter(torch.tensor(1).float())
1730
- self.scale_emb = 10
1731
- else:
1732
- self.norm_alpha = 1
1733
-
1734
  def get_input_embeddings(self):
1735
  return self.embed_tokens
1736
 
@@ -1769,7 +1198,6 @@ class MotifModel(MotifPreTrainedModel):
1769
  "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...")
1770
  use_cache = False
1771
 
1772
- # kept for BC (non `Cache` `past_key_values` inputs)
1773
  return_legacy_cache = False
1774
  if use_cache and not isinstance(past_key_values, Cache):
1775
  return_legacy_cache = True
@@ -1783,7 +1211,7 @@ class MotifModel(MotifPreTrainedModel):
1783
  "(https://huggingface.co/docs/transformers/kv_cache#legacy-cache-format)")
1784
 
1785
  if inputs_embeds is None:
1786
- inputs_embeds = self.embed_tokens(input_ids) * self.scale_emb
1787
 
1788
  if cache_position is None:
1789
  past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
@@ -1837,18 +1265,13 @@ class MotifModel(MotifPreTrainedModel):
1837
 
1838
  hidden_states = layer_outputs[0]
1839
 
1840
-
1841
- if self.use_pipeline and idx in self.split_layers:
1842
- hidden_states = torch.moreh.pipeline_assign(hidden_states)
1843
-
1844
  if use_cache:
1845
  next_decoder_cache = layer_outputs[2 if output_attentions else 1]
1846
 
1847
  if output_attentions:
1848
  all_self_attns += (layer_outputs[1], )
1849
 
1850
- # <|_2_|>
1851
- hidden_states = self.norm(hidden_states)* self.norm_alpha
1852
 
1853
  # add hidden states from the last decoder layer
1854
  if output_hidden_states:
@@ -1881,8 +1304,6 @@ class MotifModel(MotifPreTrainedModel):
1881
  output_attentions: bool,
1882
  ):
1883
  if self.config._attn_implementation == "flash_attention_2":
1884
- if MorehFlashAttention is not None:
1885
- return attention_mask
1886
  if attention_mask is not None and 0.0 in attention_mask:
1887
  return attention_mask
1888
  return None
@@ -2003,7 +1424,6 @@ class MotifModel(MotifPreTrainedModel):
2003
  return causal_mask
2004
 
2005
 
2006
- # @log_timing
2007
  class MotifForCausalLM(MotifPreTrainedModel, GenerationMixin):
2008
  _tied_weights_keys = ["lm_head.weight"]
2009
 
@@ -2013,33 +1433,19 @@ class MotifForCausalLM(MotifPreTrainedModel, GenerationMixin):
2013
  self.vocab_size = config.vocab_size
2014
  self.multi_token_heads = config.multi_token_heads
2015
 
2016
- if self.multi_token_heads is None:
2017
- self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
2018
  else:
2019
  self.tokenwise_last_layers = nn.ModuleList(
2020
  [MotifDecoderLayer(config, config.num_hidden_layers - 1) for _ in range(self.multi_token_heads)])
2021
  self.tokenwise_lm_heads = nn.ModuleList(
2022
  [nn.Linear(config.hidden_size, config.vocab_size, bias=False) for _ in range(self.multi_token_heads)])
2023
- self.should_skip_separate_backward_pass = self.multi_token_heads is not None
2024
 
2025
  # Initialize weights and apply final processing
2026
  self.post_init()
2027
-
2028
- # <|_3_|>
2029
- if config.muP:
2030
- self.lm_head.__do_scale_tager_mu_dim_base_model__=True
2031
-
2032
- # <|_4_|>
2033
- self.lm_head_alpha = 1
2034
- if config.wesar_weights:
2035
- self.lm_head_alpha = nn.Parameter(torch.tensor(1).float())
2036
-
2037
  if getattr(config, "tie_word_embeddings", True):
2038
  logger.info('tie embeddings')
2039
  self.tie_weights()
2040
- else:
2041
- # <|_5_|>
2042
- self.lm_head.__do_scale_tager_mu_dim_base_model__ = False
2043
 
2044
  def get_input_embeddings(self):
2045
  return self.model.embed_tokens
@@ -2059,101 +1465,7 @@ class MotifForCausalLM(MotifPreTrainedModel, GenerationMixin):
2059
  def get_decoder(self):
2060
  return self.model
2061
 
2062
- def multi_token_forward_backward(self,
2063
- hidden_states: torch.FloatTensor,
2064
- outputs: MotifModelOutputWithPast,
2065
- labels: torch.LongTensor,
2066
- position_ids: Optional[torch.LongTensor],
2067
- output_attentions: Optional[bool],
2068
- use_cache: Optional[bool],
2069
- cache_position: Optional[torch.LongTensor],
2070
- return_dict: Optional[bool],
2071
- num_logits_to_keep: int = 0) -> CausalLMOutputWithPast:
2072
- """
2073
- This implements the main forward-backward procedure for multi-token model training proposed in
2074
- the paper https://arxiv.org/abs/2404.19737.
2075
- Essentially,
2076
- - The multi-token model tries to predict n (instead of 1) tokens at a time.
2077
- - Applying this only during training and using first-token prediction during inference is still helpful.
2078
- - The change in architecture: when using n-token prediction, each token index (between 1 and n) has its own
2079
- (1) last attention layer and (2) lm head.
2080
- - The change in loss: sum of cross-entropy losses corresponding to each token index.
2081
- - Custom forward-backward procedure for memory efficiency: refer to the implementation of `multi_head_forward_backward`.
2082
- """
2083
- if not return_dict:
2084
- raise NotImplementedError("return_dict must be True for multi-token training")
2085
-
2086
- past_key_values = outputs.past_key_values
2087
- causal_mask = outputs.causal_mask
2088
- position_embeddings = outputs.position_embeddings
2089
-
2090
- if labels is not None:
2091
- labels = labels.to(hidden_states.device)
2092
-
2093
- def _tokenwise_forward(hidden_states: torch.Tensor, token_idx):
2094
- ## Model forward
2095
- layer = self.tokenwise_last_layers[token_idx]
2096
- lm_head = self.tokenwise_lm_heads[token_idx]
2097
-
2098
- layer_outputs = layer(
2099
- hidden_states,
2100
- attention_mask=causal_mask,
2101
- position_ids=position_ids,
2102
- past_key_values=past_key_values, # TODO: update past_key_values?
2103
- output_attentions=output_attentions,
2104
- use_cache=use_cache,
2105
- cache_position=cache_position,
2106
- position_embeddings=position_embeddings,
2107
- )
2108
- last_hidden_states = layer_outputs[0]
2109
- if num_logits_to_keep > 0:
2110
- assert labels is None
2111
- last_hidden_states = last_hidden_states[:, -num_logits_to_keep:, :]
2112
- tokenwise_logits = lm_head(last_hidden_states)
2113
-
2114
- if labels is None:
2115
- return {
2116
- "loss": None,
2117
- "logits": tokenwise_logits,
2118
- }
2119
-
2120
- ## Compute loss
2121
- shift_n = token_idx + 1
2122
- shift_logits = tokenwise_logits[..., :-shift_n, :].contiguous()
2123
- shift_labels = labels[..., shift_n:].contiguous()
2124
-
2125
- loss_fct = CrossEntropyLoss()
2126
- shift_logits = shift_logits.view(-1, self.config.vocab_size)
2127
- shift_labels = shift_labels.view(-1)
2128
-
2129
- tokenwise_loss = loss_fct(shift_logits, shift_labels)
2130
-
2131
- return {
2132
- "loss": tokenwise_loss,
2133
- "logits": tokenwise_logits,
2134
- }
2135
-
2136
- head_fns = [
2137
- lambda hidden_states, token_idx=token_idx: _tokenwise_forward(hidden_states, token_idx)
2138
- for token_idx in range(self.multi_token_heads)
2139
- ]
2140
- loss, logits = multi_head_forward_backward(hidden_states,
2141
- head_fns,
2142
- return_keys=("loss", "logits"),
2143
- return_only_first_head=True)
2144
-
2145
- if not return_dict:
2146
- output = (logits, ) + outputs[1:]
2147
- return (loss, ) + output
2148
-
2149
- return CausalLMOutputWithPast(
2150
- loss=loss,
2151
- logits=logits,
2152
- past_key_values=outputs.past_key_values,
2153
- hidden_states=outputs.hidden_states,
2154
- attentions=outputs.attentions,
2155
- )
2156
-
2157
  @add_start_docstrings_to_model_forward(MOTIF_INPUTS_DOCSTRING)
2158
  @replace_return_docstrings(output_type=CausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)
2159
  def forward(
@@ -2191,8 +1503,8 @@ class MotifForCausalLM(MotifPreTrainedModel, GenerationMixin):
2191
  ```python
2192
  >>> from transformers import AutoTokenizer, MotifForCausalLM
2193
 
2194
- >>> model = MotifForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS)
2195
- >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER)
2196
 
2197
  >>> prompt = "Hey, are you conscious? Can you talk to me?"
2198
  >>> inputs = tokenizer(prompt, return_tensors="pt")
@@ -2209,8 +1521,6 @@ class MotifForCausalLM(MotifPreTrainedModel, GenerationMixin):
2209
  return_dict = return_dict if return_dict is not None else self.config.use_return_dict
2210
 
2211
  # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
2212
- outputs_include_causal_mask = self.multi_token_heads is not None
2213
- outputs_include_position_embeddings = self.multi_token_heads is not None
2214
  outputs: MotifModelOutputWithPast = self.model(
2215
  input_ids=input_ids,
2216
  attention_mask=attention_mask,
@@ -2222,31 +1532,16 @@ class MotifForCausalLM(MotifPreTrainedModel, GenerationMixin):
2222
  output_hidden_states=output_hidden_states,
2223
  return_dict=return_dict,
2224
  cache_position=cache_position,
2225
- outputs_include_causal_mask=outputs_include_causal_mask,
2226
- outputs_include_position_embeddings=outputs_include_position_embeddings,
2227
  )
2228
 
2229
  hidden_states = outputs[0]
2230
 
2231
- if self.multi_token_heads is not None:
2232
- return self.multi_token_forward_backward(hidden_states,
2233
- outputs,
2234
- labels,
2235
- position_ids,
2236
- output_attentions,
2237
- use_cache,
2238
- cache_position,
2239
- return_dict,
2240
- num_logits_to_keep=num_logits_to_keep)
2241
-
2242
  # Only compute necessary logits, and do not upcast them to float if we are not computing the loss
2243
- hidden_states = hidden_states * self.lm_head_alpha
2244
  logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :])
2245
  logits = logits.float()
2246
 
2247
  loss = None
2248
  if labels is not None:
2249
- logits = logits
2250
  # Shift so that tokens < n predict n
2251
  shift_logits = logits[..., :-1, :].contiguous()
2252
  shift_labels = labels[..., 1:].contiguous()
 
1
  import math
2
+ from typing import List, Optional, Tuple, Union, Callable, Dict
3
 
4
  import torch
5
  import torch.utils.checkpoint
 
35
  if is_flash_attn_2_available():
36
  from transformers.modeling_flash_attention_utils import _flash_attention_forward
37
 
38
+ MorehRMSNorm = None
39
+ ScaledDotProductAttention = None
40
+ MorehFlashAttention = None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
41
 
42
  #_CHECKPOINT_FOR_DOC = "moreh/Motif-102B"
43
  _CONFIG_FOR_DOC = "MotifConfig"
44
 
 
 
 
45
  from transformers.activations import ACT2CLS as _ACT2CLS
46
  from transformers.activations import ClassInstantier
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
47
 
48
 
49
  class PolyNorm(torch.nn.Module):
50
  """
51
  A trainable activation function introduced in https://arxiv.org/html/2411.03884v1.
52
  The code is copied from https://github.com/BryceZhuo/PolyCom?tab=readme-ov-file/README.md,
53
+ with the change `* torch.rsqrt` => `/ torch.sqrt`
54
  """
55
 
56
  def __init__(self, eps=1e-6):
 
67
  x ** 2) + self.weight[2] * self._norm(x) + self.bias
68
 
69
 
70
+ CUSTOM_ACT2CLS = {"poly_norm": PolyNorm}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
71
  ACT2CLS = {**_ACT2CLS, **CUSTOM_ACT2CLS}
72
  ACT2FN = ClassInstantier(ACT2CLS)
73
 
74
 
 
75
  class MotifRMSNorm(nn.Module):
76
 
77
  def __init__(self, hidden_size, eps=1e-6):
 
84
 
85
  def forward(self, hidden_states):
86
  input_dtype = hidden_states.dtype
87
+ variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True)
 
88
  hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
89
  return self.weight * hidden_states.to(input_dtype)
90
 
 
92
  return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"
93
 
94
 
95
+ ALL_LAYERNORM_LAYERS.append(MotifRMSNorm)
96
 
97
 
98
  class MotifRotaryEmbeddingWithCache(nn.Module):
 
150
  )
151
 
152
 
 
153
  class MotifRotaryEmbedding(nn.Module):
154
 
155
  def __init__(
 
163
  config: Optional[MotifConfig] = None,
164
  ):
165
  super().__init__()
 
166
  self.rope_kwargs = {}
167
  if config is None:
168
  logger.warning_once(
 
207
  device,
208
  seq_len=seq_len,
209
  **self.rope_kwargs)
210
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
211
  self.max_seq_len_cached = seq_len
212
 
213
  if seq_len < self.original_max_seq_len and self.max_seq_len_cached > self.original_max_seq_len: # reset
 
256
  return rotated_tensor
257
 
258
 
 
259
  def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1, fused_rope=True):
260
  """
261
  Applies rotary position embeddings to the input tensors.
 
281
  q_embed = (q * cos) + (rotate_half(q) * sin)
282
  k_embed = (k * cos) + (rotate_half(k) * sin)
283
  '''
284
+
 
 
 
 
 
285
  q = q.transpose(1, 2)
286
  k = k.transpose(1, 2)
287
 
 
299
  return q_embed, k_embed
300
 
301
 
 
302
  class MotifMLP(nn.Module):
303
+
304
  def __init__(self, config):
305
  super().__init__()
306
  self.hidden_size = config.hidden_size
 
310
  self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
311
  self.act_fn = ACT2FN[config.hidden_act]
312
 
 
 
 
 
 
 
 
 
 
 
 
 
 
313
  def forward(self, hidden_state):
314
+ return self.down_proj(self.act_fn(self.gate_proj(hidden_state)) * self.up_proj(hidden_state))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
315
 
316
 
317
  def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
 
 
 
 
 
 
 
 
 
 
 
 
 
318
  return torch.repeat_interleave(hidden_states, dim=1, repeats=n_rep)
319
 
320
 
 
849
  }
850
 
851
 
 
852
  class MotifDecoderLayer(nn.Module):
853
 
854
  def __init__(self, config: MotifConfig, moe_layer: bool, layer_idx: int):
 
1119
  """
1120
 
1121
 
 
1122
  @add_start_docstrings(
1123
  "The bare Motif Model outputting raw hidden-states without any specific head on top.",
1124
  MOTIF_START_DOCSTRING,
 
1138
  self.multi_token_heads = config.multi_token_heads
1139
 
1140
  self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
 
 
1141
 
1142
  num_hidden_layers = config.num_hidden_layers if self.multi_token_heads is None else config.num_hidden_layers - 1
1143
+
 
 
 
1144
  logger.info(f'current_moe layer { moe_layer }')
1145
+ self.layers = nn.ModuleList([
1146
+ MotifDecoderLayer(config = config, layer_idx=layer_idx) for layer_idx in range(num_hidden_layers)
1147
+ ])
1148
  self._attn_implementation = config._attn_implementation
1149
  RMSNorm = MorehRMSNorm if MorehRMSNorm is not None else MotifRMSNorm
1150
  self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
 
1160
  self.gradient_checkpointing = False
1161
  self.post_init()
1162
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1163
  def get_input_embeddings(self):
1164
  return self.embed_tokens
1165
 
 
1198
  "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...")
1199
  use_cache = False
1200
 
 
1201
  return_legacy_cache = False
1202
  if use_cache and not isinstance(past_key_values, Cache):
1203
  return_legacy_cache = True
 
1211
  "(https://huggingface.co/docs/transformers/kv_cache#legacy-cache-format)")
1212
 
1213
  if inputs_embeds is None:
1214
+ inputs_embeds = self.embed_tokens(input_ids)
1215
 
1216
  if cache_position is None:
1217
  past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
 
1265
 
1266
  hidden_states = layer_outputs[0]
1267
 
 
 
 
 
1268
  if use_cache:
1269
  next_decoder_cache = layer_outputs[2 if output_attentions else 1]
1270
 
1271
  if output_attentions:
1272
  all_self_attns += (layer_outputs[1], )
1273
 
1274
+ hidden_states = self.norm(hidden_states)
 
1275
 
1276
  # add hidden states from the last decoder layer
1277
  if output_hidden_states:
 
1304
  output_attentions: bool,
1305
  ):
1306
  if self.config._attn_implementation == "flash_attention_2":
 
 
1307
  if attention_mask is not None and 0.0 in attention_mask:
1308
  return attention_mask
1309
  return None
 
1424
  return causal_mask
1425
 
1426
 
 
1427
  class MotifForCausalLM(MotifPreTrainedModel, GenerationMixin):
1428
  _tied_weights_keys = ["lm_head.weight"]
1429
 
 
1433
  self.vocab_size = config.vocab_size
1434
  self.multi_token_heads = config.multi_token_heads
1435
 
1436
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
 
1437
  else:
1438
  self.tokenwise_last_layers = nn.ModuleList(
1439
  [MotifDecoderLayer(config, config.num_hidden_layers - 1) for _ in range(self.multi_token_heads)])
1440
  self.tokenwise_lm_heads = nn.ModuleList(
1441
  [nn.Linear(config.hidden_size, config.vocab_size, bias=False) for _ in range(self.multi_token_heads)])
 
1442
 
1443
  # Initialize weights and apply final processing
1444
  self.post_init()
1445
+
 
 
 
 
 
 
 
 
 
1446
  if getattr(config, "tie_word_embeddings", True):
1447
  logger.info('tie embeddings')
1448
  self.tie_weights()
 
 
 
1449
 
1450
  def get_input_embeddings(self):
1451
  return self.model.embed_tokens
 
1465
  def get_decoder(self):
1466
  return self.model
1467
 
1468
+
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1469
  @add_start_docstrings_to_model_forward(MOTIF_INPUTS_DOCSTRING)
1470
  @replace_return_docstrings(output_type=CausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)
1471
  def forward(
 
1503
  ```python
1504
  >>> from transformers import AutoTokenizer, MotifForCausalLM
1505
 
1506
+ >>> model = MotifForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS, trust_remote_code = True)
1507
+ >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER, trust_remote_code = True)
1508
 
1509
  >>> prompt = "Hey, are you conscious? Can you talk to me?"
1510
  >>> inputs = tokenizer(prompt, return_tensors="pt")
 
1521
  return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1522
 
1523
  # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
 
 
1524
  outputs: MotifModelOutputWithPast = self.model(
1525
  input_ids=input_ids,
1526
  attention_mask=attention_mask,
 
1532
  output_hidden_states=output_hidden_states,
1533
  return_dict=return_dict,
1534
  cache_position=cache_position,
 
 
1535
  )
1536
 
1537
  hidden_states = outputs[0]
1538
 
 
 
 
 
 
 
 
 
 
 
 
1539
  # Only compute necessary logits, and do not upcast them to float if we are not computing the loss
 
1540
  logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :])
1541
  logits = logits.float()
1542
 
1543
  loss = None
1544
  if labels is not None:
 
1545
  # Shift so that tokens < n predict n
1546
  shift_logits = logits[..., :-1, :].contiguous()
1547
  shift_labels = labels[..., 1:].contiguous()