import math import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel, GenerationMixin from transformers.modeling_outputs import CausalLMOutputWithPast from .configuration_hca import HCAConfig class RMSNorm(nn.Module): def __init__(self, dim, eps=1e-5): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight class SwiGLUMlp(nn.Module): def __init__(self, d_model, d_ff): super().__init__() self.w1 = nn.Linear(d_model, d_ff, bias=False) self.w2 = nn.Linear(d_model, d_ff, bias=False) self.w3 = nn.Linear(d_ff, d_model, bias=False) def forward(self, x): return self.w3(F.silu(self.w1(x)) * self.w2(x)) def pad_to_multiple(x, multiple): T = x.shape[1] remainder = T % multiple if remainder != 0: pad_len = multiple - remainder return F.pad(x, (0, 0, 0, pad_len)), pad_len return x, 0 class HCALayer(nn.Module): def __init__(self, config: HCAConfig): super().__init__() self.d_model = config.d_model self.n_heads = config.n_heads self.d_k = config.d_model // config.n_heads self.chunk_size = config.chunk_size self.s = config.s_landmarks self.qkv_proj = nn.Linear(config.d_model, 3 * config.d_model, bias=False) self.out_proj = nn.Linear(config.d_model, config.d_model, bias=False) self.landmark_proj = nn.Parameter(torch.randn(config.chunk_size, config.s_landmarks) * (1.0 / math.sqrt(config.chunk_size))) def get_mask(self, N_c, device): C, S = self.chunk_size, self.s K_max = N_c * S + C mask = torch.zeros((N_c, 1, C, K_max), dtype=torch.bool, device=device) for c in range(N_c): past_landmarks = c * S if past_landmarks > 0: mask[c, 0, :, :past_landmarks] = True mask[c, 0, :, N_c * S : N_c * S + C] = torch.tril(torch.ones(C, C, dtype=torch.bool, device=device)) return mask def forward(self, x): B, orig_T, D = x.shape x_pad, pad_len = pad_to_multiple(x, self.chunk_size) T_pad = x_pad.shape[1] N_c = T_pad // self.chunk_size H, d_k, C, S = self.n_heads, self.d_k, self.chunk_size, self.s q, k, v = self.qkv_proj(x_pad).chunk(3, dim=-1) K_blk = k.view(B, N_c, C, H, d_k).permute(0, 3, 1, 2, 4) V_blk = v.view(B, N_c, C, H, d_k).permute(0, 3, 1, 2, 4) Q_blk = q.view(B, N_c, C, H, d_k).permute(0, 3, 1, 2, 4) K_land = torch.einsum('b h n c d, c s -> b h n s d', K_blk, self.landmark_proj).reshape(B, H, N_c * S, d_k) V_land = torch.einsum('b h n c d, c s -> b h n s d', V_blk, self.landmark_proj).reshape(B, H, N_c * S, d_k) K_all = torch.cat([K_land.unsqueeze(2).expand(B, H, N_c, N_c * S, d_k), K_blk], dim=3) V_all = torch.cat([V_land.unsqueeze(2).expand(B, H, N_c, N_c * S, d_k), V_blk], dim=3) Q_b = Q_blk.permute(0, 2, 1, 3, 4).reshape(B * N_c, H, C, d_k) K_b = K_all.permute(0, 2, 1, 3, 4).reshape(B * N_c, H, -1, d_k) V_b = V_all.permute(0, 2, 1, 3, 4).reshape(B * N_c, H, -1, d_k) mask = self.get_mask(N_c, x.device).repeat(B, 1, 1, 1) out = F.scaled_dot_product_attention(Q_b, K_b, V_b, attn_mask=mask) out = out.view(B, N_c, H, C, d_k).permute(0, 1, 3, 2, 4).reshape(B, T_pad, D) out = self.out_proj(out) return out[:, :orig_T, :] if pad_len > 0 else out def step(self, x_tok, state): B, _, D = x_tok.shape H, d_k, C, S = self.n_heads, self.d_k, self.chunk_size, self.s q, k, v = self.qkv_proj(x_tok).chunk(3, dim=-1) q = q.view(B, 1, H, d_k).transpose(1, 2) k = k.view(B, 1, H, d_k).transpose(1, 2) v = v.view(B, 1, H, d_k).transpose(1, 2) if state is None: past_K_land = torch.empty(B, H, 0, d_k, device=x_tok.device, dtype=x_tok.dtype) past_V_land = torch.empty(B, H, 0, d_k, device=x_tok.device, dtype=x_tok.dtype) curr_K, curr_V = [k], [v] else: past_K_land, past_V_land, curr_K, curr_V = state curr_K.append(k) curr_V.append(v) curr_K_tensor = torch.cat(curr_K, dim=2) curr_V_tensor = torch.cat(curr_V, dim=2) if past_K_land.shape[2] > 0: total_K = torch.cat([past_K_land, curr_K_tensor], dim=2) total_V = torch.cat([past_V_land, curr_V_tensor], dim=2) else: total_K = curr_K_tensor total_V = curr_V_tensor out = F.scaled_dot_product_attention(q, total_K, total_V) y = self.out_proj(out.transpose(1, 2).contiguous().view(B, 1, D)) if len(curr_K) >= C: K_chunk = torch.cat(curr_K, dim=2) V_chunk = torch.cat(curr_V, dim=2) K_new_land = torch.einsum('b h c d, c s -> b h s d', K_chunk, self.landmark_proj) V_new_land = torch.einsum('b h c d, c s -> b h s d', V_chunk, self.landmark_proj) past_K_land = torch.cat([past_K_land, K_new_land], dim=2) past_V_land = torch.cat([past_V_land, V_new_land], dim=2) curr_K, curr_V = [], [] return y, (past_K_land, past_V_land, curr_K, curr_V) class HCAPreTrainedModel(PreTrainedModel): config_class = HCAConfig base_model_prefix = "hca" supports_gradient_checkpointing = True _supports_cache_class = False @classmethod def _supports_default_dynamic_cache(cls) -> bool: return False class HCAForCausalLM(HCAPreTrainedModel, GenerationMixin): _supports_cache_class = False @classmethod def _supports_default_dynamic_cache(cls) -> bool: return False def __init__(self, config: HCAConfig): super().__init__(config) self.config = config self.emb = nn.Embedding(config.vocab_size, config.d_model) self.layers = nn.ModuleList([ nn.ModuleDict({ "norm1": RMSNorm(config.d_model), "op": HCALayer(config), "norm2": RMSNorm(config.d_model), "mlp": SwiGLUMlp(config.d_model, d_ff=config.d_ff) }) for _ in range(config.n_layers) ]) self.ln_f = RMSNorm(config.d_model) self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) self.post_init() def get_input_embeddings(self): return self.emb def set_input_embeddings(self, value): self.emb = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def prepare_inputs_for_generation(self, input_ids, past_key_values=None, attention_mask=None, **kwargs): if past_key_values is not None: input_ids = input_ids[:, -1:] return { "input_ids": input_ids, "past_key_values": past_key_values, "attention_mask": attention_mask } def _reorder_cache(self, past_key_values, beam_idx): if past_key_values is None: return None reordered = [] for layer_state in past_key_values: past_K_land, past_V_land, curr_K, curr_V = layer_state p_k = past_K_land.index_select(0, beam_idx) if past_K_land.numel() > 0 else past_K_land p_v = past_V_land.index_select(0, beam_idx) if past_V_land.numel() > 0 else past_V_land c_k = [k.index_select(0, beam_idx) for k in curr_K] c_v = [v.index_select(0, beam_idx) for v in curr_V] reordered.append((p_k, p_v, c_k, c_v)) return reordered def forward( self, input_ids=None, attention_mask=None, past_key_values=None, labels=None, use_cache=None, **kwargs ): B, T = input_ids.shape new_past_key_values = [] if past_key_values is None: x = self.emb(input_ids) for layer in self.layers: x_norm = layer["norm1"](x) x = x + layer["op"](x_norm) x = x + layer["mlp"](layer["norm2"](x)) if not self.training: C, S = self.config.chunk_size, self.config.s_landmarks N_full = T // C leftover = T % C x_init = self.emb(input_ids) for layer in self.layers: op = layer["op"] x_n = layer["norm1"](x_init) _, k, v = op.qkv_proj(x_n).chunk(3, dim=-1) k = k.view(B, T, op.n_heads, op.d_k).transpose(1, 2) v = v.view(B, T, op.n_heads, op.d_k).transpose(1, 2) if N_full > 0: k_full = k[:, :, :N_full * C, :].reshape(B, op.n_heads, N_full, C, op.d_k) v_full = v[:, :, :N_full * C, :].reshape(B, op.n_heads, N_full, C, op.d_k) k_land = torch.einsum('b h n c d, c s -> b h n s d', k_full, op.landmark_proj).reshape(B, op.n_heads, N_full * S, op.d_k) v_land = torch.einsum('b h n c d, c s -> b h n s d', v_full, op.landmark_proj).reshape(B, op.n_heads, N_full * S, op.d_k) else: k_land = torch.empty(B, op.n_heads, 0, op.d_k, device=input_ids.device, dtype=k.dtype) v_land = torch.empty(B, op.n_heads, 0, op.d_k, device=input_ids.device, dtype=v.dtype) curr_K = [k[:, :, N_full * C + i : N_full * C + i + 1, :] for i in range(leftover)] curr_V = [v[:, :, N_full * C + i : N_full * C + i + 1, :] for i in range(leftover)] new_past_key_values.append((k_land, v_land, curr_K, curr_V)) x_init = x_init + layer["op"](x_n) x_init = x_init + layer["mlp"](layer["norm2"](x_init)) else: new_past_key_values = None else: x = self.emb(input_ids) for idx, layer in enumerate(self.layers): x_norm = layer["norm1"](x) y, next_state = layer["op"].step(x_norm, past_key_values[idx]) x = x + y x = x + layer["mlp"](layer["norm2"](x)) new_past_key_values.append(next_state) logits = self.lm_head(self.ln_f(x)) loss = None if labels is not None: loss = F.cross_entropy(logits.view(-1, self.config.vocab_size), labels.view(-1), ignore_index=-100) return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=new_past_key_values if not self.training else None )