bilzepython/minimind-3v
09
1import math, torch, torch.nn.functional as F2from torch import nn3from transformers.activations import ACT2FN4from transformers import PreTrainedModel, GenerationMixin, PretrainedConfig5from transformers.modeling_outputs import MoeCausalLMOutputWithPast6 7# ๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐8# MiniMind Config9# ๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐10class MiniMindConfig(PretrainedConfig):11 model_type = "minimind"12 def __init__(self, hidden_size=768, num_hidden_layers=8, use_moe=False, **kwargs):13 super().__init__(**kwargs)14 self.hidden_size = hidden_size15 self.num_hidden_layers = num_hidden_layers16 self.use_moe = use_moe17 self.dropout = kwargs.get("dropout", 0.0)18 self.vocab_size = kwargs.get("vocab_size", 6400)19 self.bos_token_id = kwargs.get("bos_token_id", 1)20 self.eos_token_id = kwargs.get("eos_token_id", 2)21 self.flash_attn = kwargs.get("flash_attn", True)22 self.num_attention_heads = kwargs.get("num_attention_heads", 8)23 self.num_key_value_heads = kwargs.get("num_key_value_heads", 4)24 self.head_dim = kwargs.get("head_dim", self.hidden_size // self.num_attention_heads)25 self.hidden_act = kwargs.get("hidden_act", 'silu')26 self.intermediate_size = kwargs.get("intermediate_size", math.ceil(hidden_size * math.pi / 64) * 64)27 self.max_position_embeddings = kwargs.get("max_position_embeddings", 32768)28 self.rms_norm_eps = kwargs.get("rms_norm_eps", 1e-6)29 self.rope_theta = kwargs.get("rope_theta", 1e6)30 self.inference_rope_scaling = kwargs.get("inference_rope_scaling", False)31 self.rope_scaling = {32 "beta_fast": 32,33 "beta_slow": 1,34 "factor": 16,35 "original_max_position_embeddings": 2048,36 "attention_factor": 1.0,37 "type": "yarn"38 } if self.inference_rope_scaling else None39 ### MoE specific configs (ignored if use_moe = False)40 self.num_experts = kwargs.get("num_experts", 4)41 self.num_experts_per_tok = kwargs.get("num_experts_per_tok", 1)42 self.moe_intermediate_size = kwargs.get("moe_intermediate_size", self.intermediate_size)43 self.norm_topk_prob = kwargs.get("norm_topk_prob", True)44 self.router_aux_loss_coef = kwargs.get("router_aux_loss_coef", 5e-4)45 46# ๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐47# MiniMind Model48# ๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐๐49class RMSNorm(torch.nn.Module):50 def __init__(self, dim: int, eps: float = 1e-5):51 super().__init__()52 self.eps = eps53 self.weight = nn.Parameter(torch.ones(dim))54 55 def norm(self, x):56 return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)57 58 def forward(self, x):59 return (self.weight * self.norm(x.float())).type_as(x)60 61def precompute_freqs_cis(dim: int, end: int = int(32 * 1024), rope_base: float = 1e6, rope_scaling: dict = None):62 freqs, attn_factor = 1.0 / (rope_base ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)), 1.063 if rope_scaling is not None: # YaRN: f'(i) = f(i)((1-ฮณ) + ฮณ/s), where ฮณโ[0,1] is linear ramp64 orig_max, factor, beta_fast, beta_slow, attn_factor = (65 rope_scaling.get("original_max_position_embeddings", 2048), rope_scaling.get("factor", 16),66 rope_scaling.get("beta_fast", 32.0), rope_scaling.get("beta_slow", 1.0), rope_scaling.get("attention_factor", 1.0)67 )68 if end / orig_max > 1.0:69 inv_dim = lambda b: (dim * math.log(orig_max / (b * 2 * math.pi))) / (2 * math.log(rope_base))70 low, high = max(math.floor(inv_dim(beta_fast)), 0), min(math.ceil(inv_dim(beta_slow)), dim // 2 - 1)71 ramp = torch.clamp((torch.arange(dim // 2, device=freqs.device).float() - low) / max(high - low, 0.001), 0, 1)72 freqs = freqs * (1 - ramp + ramp / factor)73 t = torch.arange(end, device=freqs.device)74 freqs = torch.outer(t, freqs).float()75 freqs_cos = torch.cat([torch.cos(freqs), torch.cos(freqs)], dim=-1) * attn_factor76 freqs_sin = torch.cat([torch.sin(freqs), torch.sin(freqs)], dim=-1) * attn_factor77 return freqs_cos, freqs_sin78 79def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1):80 def rotate_half(x): return torch.cat((-x[..., x.shape[-1] // 2:], x[..., : x.shape[-1] // 2]), dim=-1)81 q_embed = ((q * cos.unsqueeze(unsqueeze_dim)) + (rotate_half(q) * sin.unsqueeze(unsqueeze_dim))).to(q.dtype)82 k_embed = ((k * cos.unsqueeze(unsqueeze_dim)) + (rotate_half(k) * sin.unsqueeze(unsqueeze_dim))).to(k.dtype)83 return q_embed, k_embed84 85def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:86 bs, slen, num_key_value_heads, head_dim = x.shape87 if n_rep == 1: return x88 return (x[:, :, :, None, :].expand(bs, slen, num_key_value_heads, n_rep, head_dim).reshape(bs, slen, num_key_value_heads * n_rep, head_dim))89 90class Attention(nn.Module):91 def __init__(self, config: MiniMindConfig):92 super().__init__()93 self.num_key_value_heads = config.num_attention_heads if config.num_key_value_heads is None else config.num_key_value_heads94 self.n_local_heads = config.num_attention_heads95 self.n_local_kv_heads = self.num_key_value_heads96 self.n_rep = self.n_local_heads // self.n_local_kv_heads97 self.head_dim = config.head_dim98 self.q_proj = nn.Linear(config.hidden_size, config.num_attention_heads * self.head_dim, bias=False)99 self.k_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)100 self.v_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)101 self.o_proj = nn.Linear(config.num_attention_heads * self.head_dim, config.hidden_size, bias=False)102 self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)103 self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)104 self.attn_dropout = nn.Dropout(config.dropout)105 self.resid_dropout = nn.Dropout(config.dropout)106 self.dropout = config.dropout107 self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention') and config.flash_attn108 109 def forward(self, x, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None):110 bsz, seq_len, _ = x.shape111 xq, xk, xv = self.q_proj(x), self.k_proj(x), self.v_proj(x)112 xq = xq.view(bsz, seq_len, self.n_local_heads, self.head_dim)113 xk = xk.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)114 xv = xv.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)115 xq, xk = self.q_norm(xq), self.k_norm(xk)116 cos, sin = position_embeddings117 xq, xk = apply_rotary_pos_emb(xq, xk, cos, sin)118 if past_key_value is not None:119 xk = torch.cat([past_key_value[0], xk], dim=1)120 xv = torch.cat([past_key_value[1], xv], dim=1)121 past_kv = (xk, xv) if use_cache else None122 xq, xk, xv = (xq.transpose(1, 2), repeat_kv(xk, self.n_rep).transpose(1, 2), repeat_kv(xv, self.n_rep).transpose(1, 2))123 if self.flash and (seq_len > 1) and (past_key_value is None) and (attention_mask is None or torch.all(attention_mask == 1)):124 output = F.scaled_dot_product_attention(xq, xk, xv, dropout_p=self.dropout if self.training else 0.0, is_causal=True)125 else:126 scores = (xq @ xk.transpose(-2, -1)) / math.sqrt(self.head_dim)127 scores[:, :, :, -seq_len:] += torch.full((seq_len, seq_len), float("-inf"), device=scores.device).triu(1)128 if attention_mask is not None: scores += (1.0 - attention_mask.unsqueeze(1).unsqueeze(2)) * -1e9129 output = self.attn_dropout(F.softmax(scores.float(), dim=-1).type_as(xq)) @ xv130 output = output.transpose(1, 2).reshape(bsz, seq_len, -1)131 output = self.resid_dropout(self.o_proj(output))132 return output, past_kv133 134class FeedForward(nn.Module):135 def __init__(self, config: MiniMindConfig, intermediate_size: int = None):136 super().__init__()137 intermediate_size = intermediate_size or config.intermediate_size138 self.gate_proj = nn.Linear(config.hidden_size, intermediate_size, bias=False)139 self.down_proj = nn.Linear(intermediate_size, config.hidden_size, bias=False)140 self.up_proj = nn.Linear(config.hidden_size, intermediate_size, bias=False)141 self.act_fn = ACT2FN[config.hidden_act]142 143 def forward(self, x):144 return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))145 146class MOEFeedForward(nn.Module):147 def __init__(self, config: MiniMindConfig):148 super().__init__()149 self.config = config150 self.gate = nn.Linear(config.hidden_size, config.num_experts, bias=False)151 self.experts = nn.ModuleList([FeedForward(config, intermediate_size=config.moe_intermediate_size) for _ in range(config.num_experts)])152 self.act_fn = ACT2FN[config.hidden_act]153 154 def forward(self, x):155 batch_size, seq_len, hidden_dim = x.shape156 x_flat = x.view(-1, hidden_dim)157 scores = F.softmax(self.gate(x_flat), dim=-1)158 topk_weight, topk_idx = torch.topk(scores, k=self.config.num_experts_per_tok, dim=-1, sorted=False)159 if self.config.norm_topk_prob: topk_weight = topk_weight / (topk_weight.sum(dim=-1, keepdim=True) + 1e-20)160 y = torch.zeros_like(x_flat)161 for i, expert in enumerate(self.experts):162 mask = (topk_idx == i)163 if mask.any():164 token_idx = mask.any(dim=-1).nonzero().flatten()165 weight = topk_weight[mask].view(-1, 1)166 y.index_add_(0, token_idx, (expert(x_flat[token_idx]) * weight).to(y.dtype))167 elif self.training:168 y[0, 0] += 0 * sum(p.sum() for p in expert.parameters())169 if self.training and self.config.router_aux_loss_coef > 0:170 load = F.one_hot(topk_idx, self.config.num_experts).float().mean(0)171 self.aux_loss = (load * scores.mean(0)).sum() * self.config.num_experts * self.config.router_aux_loss_coef172 else:173 self.aux_loss = scores.new_zeros(1).squeeze()174 return y.view(batch_size, seq_len, hidden_dim)175 176class MiniMindBlock(nn.Module):177 def __init__(self, layer_id: int, config: MiniMindConfig):178 super().__init__()179 self.self_attn = Attention(config)180 self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)181 self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)182 self.mlp = FeedForward(config) if not config.use_moe else MOEFeedForward(config)183 184 def forward(self, hidden_states, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None):185 residual = hidden_states186 hidden_states, present_key_value = self.self_attn(187 self.input_layernorm(hidden_states), position_embeddings,188 past_key_value, use_cache, attention_mask189 )190 hidden_states += residual191 hidden_states = hidden_states + self.mlp(self.post_attention_layernorm(hidden_states))192 return hidden_states, present_key_value193 194class MiniMindModel(nn.Module):195 def __init__(self, config: MiniMindConfig):196 super().__init__()197 self.config = config198 self.vocab_size, self.num_hidden_layers = config.vocab_size, config.num_hidden_layers199 self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)200 self.dropout = nn.Dropout(config.dropout)201 self.layers = nn.ModuleList([MiniMindBlock(l, config) for l in range(self.num_hidden_layers)])202 self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)203 freqs_cos, freqs_sin = precompute_freqs_cis(dim=config.head_dim, end=config.max_position_embeddings, rope_base=config.rope_theta, rope_scaling=config.rope_scaling)204 self.register_buffer("freqs_cos", freqs_cos, persistent=False)205 self.register_buffer("freqs_sin", freqs_sin, persistent=False)206 207 def forward(self, input_ids, attention_mask=None, past_key_values=None, use_cache=False, **kwargs):208 batch_size, seq_length = input_ids.shape209 if hasattr(past_key_values, 'layers'): past_key_values = None210 past_key_values = past_key_values or [None] * len(self.layers)211 start_pos = past_key_values[0][0].shape[1] if past_key_values[0] is not None else 0212 hidden_states = self.dropout(self.embed_tokens(input_ids))213 position_embeddings = (self.freqs_cos[start_pos:start_pos + seq_length], self.freqs_sin[start_pos:start_pos + seq_length])214 presents = []215 for layer, past_key_value in zip(self.layers, past_key_values):216 hidden_states, present = layer(217 hidden_states,218 position_embeddings,219 past_key_value=past_key_value,220 use_cache=use_cache,221 attention_mask=attention_mask222 )223 presents.append(present)224 hidden_states = self.norm(hidden_states)225 aux_loss = sum([l.mlp.aux_loss for l in self.layers if isinstance(l.mlp, MOEFeedForward)], hidden_states.new_zeros(1).squeeze())226 return hidden_states, presents, aux_loss227 228class MiniMindForCausalLM(PreTrainedModel, GenerationMixin):229 config_class = MiniMindConfig230 def __init__(self, config: MiniMindConfig = None):231 self.config = config or MiniMindConfig()232 super().__init__(self.config)233 self.model = MiniMindModel(self.config)234 self.lm_head = nn.Linear(self.config.hidden_size, self.config.vocab_size, bias=False)235 self.model.embed_tokens.weight = self.lm_head.weight236 237 def forward(self, input_ids, attention_mask=None, past_key_values=None, use_cache=False, logits_to_keep=0, labels=None, **kwargs):238 hidden_states, past_key_values, aux_loss = self.model(input_ids, attention_mask, past_key_values, use_cache, **kwargs)239 slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep240 logits = self.lm_head(hidden_states[:, slice_indices, :])241 loss = None242 if labels is not None:243 x, y = logits[..., :-1, :].contiguous(), labels[..., 1:].contiguous()244 loss = F.cross_entropy(x.view(-1, x.size(-1)), y.view(-1), ignore_index=-100)245 return MoeCausalLMOutputWithPast(loss=loss, aux_loss=aux_loss, logits=logits, past_key_values=past_key_values, hidden_states=hidden_states)246 247 # https://github.com/jingyaogong/minimind/discussions/611248 @torch.inference_mode()249 def generate(self, inputs=None, attention_mask=None, max_new_tokens=8192, temperature=0.85, top_p=0.85, top_k=50, eos_token_id=2, streamer=None, use_cache=True, num_return_sequences=1, do_sample=True, repetition_penalty=1.0, **kwargs):250 input_ids = kwargs.pop("input_ids", inputs).repeat(num_return_sequences, 1)251 attention_mask = attention_mask.repeat(num_return_sequences, 1) if attention_mask is not None else None252 past_key_values = kwargs.pop("past_key_values", None)253 finished = torch.zeros(input_ids.shape[0], dtype=torch.bool, device=input_ids.device)254 if streamer: streamer.put(input_ids.cpu())255 for _ in range(max_new_tokens):256 past_len = past_key_values[0][0].shape[1] if past_key_values else 0257 outputs = self.forward(input_ids[:, past_len:], attention_mask, past_key_values, use_cache=use_cache, **kwargs)258 attention_mask = torch.cat([attention_mask, attention_mask.new_ones(attention_mask.shape[0], 1)], -1) if attention_mask is not None else None259 logits = outputs.logits[:, -1, :] / temperature260 if repetition_penalty != 1.0:261 for i in range(input_ids.shape[0]): logits[i, torch.unique(input_ids[i])] /= repetition_penalty262 if top_k > 0: 263 logits[logits < torch.topk(logits, top_k)[0][..., -1, None]] = -float('inf')264 if top_p < 1.0:265 sorted_logits, sorted_indices = torch.sort(logits, descending=True)266 mask = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1) > top_p267 mask[..., 1:], mask[..., 0] = mask[..., :-1].clone(), 0268 logits[mask.scatter(1, sorted_indices, mask)] = -float('inf')269 next_token = torch.multinomial(torch.softmax(logits, dim=-1), num_samples=1) if do_sample else torch.argmax(logits, dim=-1, keepdim=True)270 if eos_token_id is not None: next_token = torch.where(finished.unsqueeze(-1), next_token.new_full((next_token.shape[0], 1), eos_token_id), next_token)271 input_ids = torch.cat([input_ids, next_token], dim=-1)272 past_key_values = outputs.past_key_values if use_cache else None273 if streamer: streamer.put(next_token.cpu())274 if eos_token_id is not None:275 finished |= next_token.squeeze(-1).eq(eos_token_id)276 if finished.all(): break277 if streamer: streamer.end()278 if kwargs.get("return_kv"): return {'generated_ids': input_ids, 'past_kv': past_key_values}279 return input_ids