class CausalSelfAttention(nn.Module):
def __init__(self, d_model: int, n_heads: int) -> None:
super().__init__()
if d_model % n_heads != 0:
raise ValueError("d_model doit être divisible par n_heads")
self.n_heads = n_heads
self.head_dim = d_model // n_heads
self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
self.out = nn.Linear(d_model, d_model, bias=False)
def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
batch, length, _ = x.shape
return x.view(batch, length, self.n_heads, self.head_dim).transpose(1, 2)
def forward(
self,
x: torch.Tensor,
past_kv: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]:
q, k, v = self.qkv(x).chunk(3, dim=-1)
q, k, v = map(self._split_heads, (q, k, v))
past_length = 0 if past_kv is None else past_kv[0].shape[2]
if past_kv is not None:
k = torch.cat((past_kv[0], k), dim=2)
v = torch.cat((past_kv[1], v), dim=2)
query_positions = past_length + torch.arange(q.shape[2], device=x.device)
key_positions = torch.arange(k.shape[2], device=x.device)
causal = key_positions[None, :] <= query_positions[:, None]
scores = q @ k.transpose(-2, -1) / math.sqrt(self.head_dim)
scores = scores.masked_fill(~causal[None, None, :, :], float("-inf"))
context = torch.softmax(scores, dim=-1) @ v
context = context.transpose(1, 2).contiguous().view(x.shape)
return self.out(context), (k, v)
class TransformerBlock(nn.Module):
def __init__(self, d_model: int, n_heads: int, expansion: int = 4) -> None:
super().__init__()
self.norm_attention = nn.LayerNorm(d_model)
self.attention = CausalSelfAttention(d_model, n_heads)
self.norm_mlp = nn.LayerNorm(d_model)
self.mlp = nn.Sequential(
nn.Linear(d_model, expansion * d_model),
nn.GELU(),
nn.Linear(expansion * d_model, d_model),
)
def forward(self, x, past_kv=None):
attended, present_kv = self.attention(self.norm_attention(x), past_kv)
x = x + attended
x = x + self.mlp(self.norm_mlp(x))
return x, present_kv
class TinyCausalLM(nn.Module):
def __init__(
self,
vocab_size: int = 256,
d_model: int = 160,
n_heads: int = 5,
n_layers: int = 3,
max_length: int = 512,
) -> None:
super().__init__()
self.token_embedding = nn.Embedding(vocab_size, d_model)
self.position_embedding = nn.Embedding(max_length, d_model)
self.blocks = nn.ModuleList(
[TransformerBlock(d_model, n_heads) for _ in range(n_layers)]
)
self.norm = nn.LayerNorm(d_model)
self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
self.n_layers = n_layers
self.n_heads = n_heads
self.head_dim = d_model // n_heads
def forward(self, token_ids, past_key_values=None):
past_length = 0 if past_key_values is None else past_key_values[0][0].shape[2]
positions = past_length + torch.arange(token_ids.shape[1], device=token_ids.device)
x = self.token_embedding(token_ids) + self.position_embedding(positions)[None, :, :]
present = []
for layer_index, block in enumerate(self.blocks):
past = None if past_key_values is None else past_key_values[layer_index]
x, layer_kv = block(x, past)
present.append(layer_kv)
return self.lm_head(self.norm(x)), tuple(present)
model = TinyCausalLM().eval()
parameter_count = sum(parameter.numel() for parameter in model.parameters())
print(f"Modèle créé : {parameter_count:,} paramètres, {model.n_layers} couches, "
f"{model.n_heads} têtes, dimension de tête={model.head_dim}")