import torch.nn as nn class GQAConfig: num_query_heads = 32 num_kv_heads = 8 head_dim = 128 hidden_size = 4096 class GQAAttention(nn.Module): def __init__(self, config): super().__init__() self.q = nn.Linear(config.hidden_size, config.num_query_heads * config.head_dim) self.k = nn.Linear(config.hidden_size, config.num_kv_heads * config.head_dim) self.v = nn.Linear(config.hidden_size, config.num_kv_heads * config.head_dim) # GQA logic: repeat KV heads to match query head count # implementation via functional repeat_kv