初始化项目,由ModelHub XC社区提供模型

Model: sabari2005/cyberslm-instruct
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-08-29 19:29:19 +08:00
commit 4244787e58
25 changed files with 282452 additions and 0 deletions

35
.gitattributes vendored Normal file
View File

@@ -0,0 +1,35 @@
*.7z filter=lfs diff=lfs merge=lfs -text
*.arrow filter=lfs diff=lfs merge=lfs -text
*.bin filter=lfs diff=lfs merge=lfs -text
*.bz2 filter=lfs diff=lfs merge=lfs -text
*.ckpt filter=lfs diff=lfs merge=lfs -text
*.ftz filter=lfs diff=lfs merge=lfs -text
*.gz filter=lfs diff=lfs merge=lfs -text
*.h5 filter=lfs diff=lfs merge=lfs -text
*.joblib filter=lfs diff=lfs merge=lfs -text
*.lfs.* filter=lfs diff=lfs merge=lfs -text
*.mlmodel filter=lfs diff=lfs merge=lfs -text
*.model filter=lfs diff=lfs merge=lfs -text
*.msgpack filter=lfs diff=lfs merge=lfs -text
*.npy filter=lfs diff=lfs merge=lfs -text
*.npz filter=lfs diff=lfs merge=lfs -text
*.onnx filter=lfs diff=lfs merge=lfs -text
*.ot filter=lfs diff=lfs merge=lfs -text
*.parquet filter=lfs diff=lfs merge=lfs -text
*.pb filter=lfs diff=lfs merge=lfs -text
*.pickle filter=lfs diff=lfs merge=lfs -text
*.pkl filter=lfs diff=lfs merge=lfs -text
*.pt filter=lfs diff=lfs merge=lfs -text
*.pth filter=lfs diff=lfs merge=lfs -text
*.rar filter=lfs diff=lfs merge=lfs -text
*.safetensors filter=lfs diff=lfs merge=lfs -text
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
*.tar.* filter=lfs diff=lfs merge=lfs -text
*.tar filter=lfs diff=lfs merge=lfs -text
*.tflite filter=lfs diff=lfs merge=lfs -text
*.tgz filter=lfs diff=lfs merge=lfs -text
*.wasm filter=lfs diff=lfs merge=lfs -text
*.xz filter=lfs diff=lfs merge=lfs -text
*.zip filter=lfs diff=lfs merge=lfs -text
*.zst filter=lfs diff=lfs merge=lfs -text
*tfevents* filter=lfs diff=lfs merge=lfs -text

182
README.md Normal file
View File

@@ -0,0 +1,182 @@
---
license: apache-2.0
language:
- en
library_name: pytorch
pipeline_tag: text-generation
base_model: sabari2005/cyberslm-base
tags:
- cybersecurity
- security
- small-language-model
- instruction-tuned
- sft
- causal-lm
---
# CyberSLM-instruct — 33.5M-parameter cybersecurity assistant
Instruction-tuned from
[**sabari2005/cyberslm-base**](https://huggingface.co/sabari2005/cyberslm-base)
on 23,540 cybersecurity conversations.
**Code:** [github.com/Sabari2005/cyberslm](https://github.com/Sabari2005/cyberslm)
## Read this before using it
At 33.5M parameters this model learned the **shape** of a good answer — markdown
structure, numbered steps, worked examples, mitigation sections — and is
**frequently wrong about the content**.
Measured over 8 greedy prompts:
| category | mean 8-gram repetition | stopped on EOS |
|---|---:|---:|
| security | 18.7% | 1 / 4 |
| general | 18.9% | 1 / 2 |
| code | 19.9% | 0 / 2 |
| **overall** | **19.0%** | **2 / 8** |
**Do not use this as a factual security reference. Do not run code it
generates.** It is a demonstration of a correct end-to-end training pipeline and
a baseline for scaling.
### Where it works
Well-covered in-domain questions come out correct and terminate cleanly:
> **What is SQL injection and how do I prevent it?**
>
> SQL injection (SQLi) is a security vulnerability that allows attackers to
> manipulate database queries by injecting malicious SQL code through input
> fields. It occurs when user-supplied data is improperly sanitized, allowing
> attackers to manipulate the database. For example, if an attacker submits a
> username like `admin' --` as the username, the query becomes:
> `SELECT * FROM users WHERE username = '[input]' AND password = '[input]'`.
> **What is a firewall?**
>
> A firewall is a network security device or software that monitors and controls
> incoming and outgoing network traffic based on predefined security rules. It
> acts as a barrier between trusted internal networks and untrusted external
> networks, such as the internet.
### Where it fails
All observed, reproducible, and the majority case:
* **Wrong content, confident tone.** Asked to contrast symmetric and asymmetric
encryption it answered about hashing and IKE, never addressing the question.
* **Circular definitions.** "A buffer overflow is a type of buffer overflow that
could lead to arbitrary code execution."
* **Topically-adjacent but wrong vocabulary.** Asked how to investigate a
phishing email it produced `SameSite` and `Strict` — real security terms,
wrong topic (they are cookie attributes).
* **Degenerate loops in code.** `port: The port to use` repeated to the token
limit.
* **Unreliable termination.** Only 2 of 8 prompts stopped on EOS; the rest ran
to the token limit.
These are consequences of scale, not of the training run — the loss curve is
healthy and the pipeline is machine-verified (35 architecture checks, 173 tests).
## Model details
| | |
|---|---|
| parameters | 33,531,264 |
| architecture | 12 layers, d_model 384, 6 heads, SwiGLU 1024, RoPE, RMSNorm, tied head |
| context | 2048 |
| vocab | 32,000 (SentencePiece BPE) |
| base model | sabari2005/cyberslm-base |
| SFT data | 23,540 conversations, 15.1M supervised tokens |
| epochs | 3 (2,208 optimizer steps) |
| optimiser | AdamW, lr 2e-5, 3% warmup, cosine, bf16 |
| best val loss | 2.2627 (response tokens only) |
Loss is computed on assistant responses only; prompts are masked. 88% of tokens
in the SFT set are supervised.
## Usage
```bash
pip install torch sentencepiece
git clone https://huggingface.co/sabari2005/cyberslm-instruct
cd cyberslm-instruct
python infer_chat.py --prompt "What is SQL injection and how do I prevent it?"
```
Interactive:
```bash
python infer_chat.py --interactive
```
Options:
```bash
python infer_chat.py \
--prompt "What is a buffer overflow?" \
--max-new-tokens 200 \
--temperature 0.0 # 0 = greedy, recommended for this model
```
### Prompt format
The model was trained on this exact layout, with a real BOS token id prepended
and EOS terminating each response:
```
### User:
{question}
### Assistant:
{response}<eos>
```
**Build prompts with the bundled formatter** (`infer_chat.py` does this).
Hand-assembling the string produces different token ids at every segment
boundary, because SentencePiece prepends a word-boundary marker per `encode()`
call — the model then sees something it was never trained on.
```python
import torch
from configs.sft_config import default_config
from data.prompt_formatter import PromptFormatter, Tokenizer
from model.cyberslm import CyberSLM
cfg = default_config()
cfg.tokenizer.model_path = "tokenizer/tokenizer.model"
cfg.model.max_seq_len = cfg.data.max_seq_len = 2048
tok = Tokenizer(cfg.tokenizer.model_path)
fmt = PromptFormatter(cfg=cfg, tokenizer=tok)
model = CyberSLM(cfg.model)
model.load_state_dict(torch.load("models/instruct.pt", map_location="cpu",
weights_only=False))
model.eval()
ids = fmt.format_for_inference({"messages": [{"role": "user",
"content": "What is XSS?"}]})
out = model.generate(torch.tensor([ids]), max_new_tokens=200,
temperature=0.0, eos_id=tok.eos_id)
print(tok.decode(out[0, len(ids):].tolist()))
```
Decoding uses a KV cache — roughly 5070 tok/s on CPU.
## Intended use
Research into small language models; a scaling baseline; a demonstration of a
verified training pipeline. **Not** for security advice, incident response, code
generation, or anything where accuracy matters.
## Training data
Not published. Curated cybersecurity instruction data; not redistributed.
## License
Apache-2.0 for the code and weights. Verify licensing for downstream use against
the sources the data was curated from.

32
config.json Normal file
View File

@@ -0,0 +1,32 @@
{
"architectures": [
"LlamaForCausalLM"
],
"attention_bias": false,
"attention_dropout": 0.0,
"bos_token_id": 2,
"dtype": "float32",
"eos_token_id": 3,
"head_dim": 64,
"hidden_act": "silu",
"hidden_size": 384,
"initializer_range": 0.02,
"intermediate_size": 1024,
"max_position_embeddings": 2048,
"mlp_bias": false,
"model_type": "llama",
"num_attention_heads": 6,
"num_hidden_layers": 12,
"num_key_value_heads": 6,
"pad_token_id": 0,
"pretraining_tp": 1,
"rms_norm_eps": 1e-06,
"rope_parameters": {
"rope_theta": 10000.0,
"rope_type": "default"
},
"tie_word_embeddings": true,
"transformers_version": "5.15.1",
"use_cache": true,
"vocab_size": 32000
}

View File

@@ -0,0 +1,35 @@
"""
CyberSLM model package.
Public API (grows as phases are added):
Phase 1: CyberSLMConfig, default_config, RMSNorm, RotaryPositionEmbedding, apply_rope
Phase 2: MultiHeadSelfAttention, SwiGLUFeedForward
Phase 3: DecoderBlock, CyberSLM, build_model
Phase 4: (training engine lives in cyberslm.training)
"""
from cyberslm.model.config import CyberSLMConfig, default_config
from cyberslm.model.norm import RMSNorm
from cyberslm.model.rope import RotaryPositionEmbedding, apply_rope
from cyberslm.model.attention import MultiHeadSelfAttention
from cyberslm.model.ffn import SwiGLUFeedForward
from cyberslm.model.block import DecoderBlock
from cyberslm.model.model import CyberSLM, build_model, count_parameters, model_summary
__all__ = [
# Phase 1
"CyberSLMConfig",
"default_config",
"RMSNorm",
"RotaryPositionEmbedding",
"apply_rope",
# Phase 2
"MultiHeadSelfAttention",
"SwiGLUFeedForward",
# Phase 3
"DecoderBlock",
"CyberSLM",
"build_model",
"count_parameters",
"model_summary",
]

302
cyberslm/model/attention.py Normal file
View File

@@ -0,0 +1,302 @@
"""
Multi-Head Self Attention (MHSA)
=================================
Standard scaled dot-product multi-head self attention with:
- Rotary Position Embedding (RoPE) on queries and keys
- Causal (auto-regressive) masking
- No bias on projection layers
- Pre-norm placement handled by the enclosing DecoderBlock
Mathematical definition
-----------------------
Given input X ∈ ^{B×T×d}:
Q = X Wq, K = X Wk, V = X Wv (projections, no bias)
Split into H heads, each of dimension d_h = d / H:
Qₕ, Kₕ = RoPE(Qₕ), RoPE(Kₕ) (apply rotary embeddings)
Scaled dot-product attention per head:
Aₕ = softmax( (Qₕ Kₕᵀ) / √d_h + mask ) Vₕ
where mask[i,j] = 0 if j ≤ i else −∞ (causal constraint).
Concatenate and project:
output = concat(A₁, ..., A_H) Wo
Complexity
----------
Time : O(T² · d) — quadratic in sequence length (standard attention)
Space: O(T² · H) — attention weight matrix per head
Numerical stability
-------------------
- Scaling by 1/√d_h keeps the pre-softmax logits in a well-conditioned
range, preventing vanishing gradients from very peaked softmax outputs.
- Softmax is computed by PyTorch's numerically stable implementation
(subtract max before exp).
- RoPE is applied in float32 (see rope.py).
- Causal mask adds −∞ (not a large negative number) so masked positions
become exactly 0 after softmax — no gradient leakage.
FlashAttention compatibility
-----------------------------
The forward pass is written in a way that is structurally compatible with
a future drop-in replacement by ``torch.nn.functional.scaled_dot_product_attention``
(PyTorch 2.0+) or the ``flash-attn`` library. To migrate:
1. Replace the manual QKᵀ/softmax/V block with:
F.scaled_dot_product_attention(q, k, v, attn_mask=None,
dropout_p=0.0, is_causal=True)
2. Remove the manual mask addition (is_causal=True handles it).
"""
from __future__ import annotations
import math
from typing import Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from cyberslm.model.config import CyberSLMConfig
from cyberslm.model.rope import RotaryPositionEmbedding, apply_rope
def _causal_bias(
q_len: int,
k_len: int,
past_len: int,
dtype: torch.dtype,
device: torch.device,
) -> Tensor:
"""
Additive causal mask of shape ``(q_len, k_len)`` for a query block that
starts at absolute position ``past_len``.
Query row ``i`` represents absolute position ``past_len + i`` and may attend
to key columns ``0 .. past_len + i`` inclusive; everything after is -inf.
With ``past_len == 0`` this reduces to the usual upper-triangular mask.
"""
q_pos = torch.arange(q_len, device=device).unsqueeze(1) + past_len # (q,1)
k_pos = torch.arange(k_len, device=device).unsqueeze(0) # (1,k)
return torch.where(
k_pos <= q_pos,
torch.zeros((), dtype=dtype, device=device),
torch.full((), float("-inf"), dtype=dtype, device=device),
)
class MultiHeadSelfAttention(nn.Module):
"""
Multi-Head Self Attention with RoPE and causal masking.
This module owns the four projection matrices (Wq, Wk, Wv, Wo),
the RoPE cache, and the causal mask buffer.
Parameters
----------
config : CyberSLMConfig
Validated model configuration.
Attributes
----------
q_proj : nn.Linear ``(hidden_dim, hidden_dim)``, no bias
k_proj : nn.Linear ``(hidden_dim, hidden_dim)``, no bias
v_proj : nn.Linear ``(hidden_dim, hidden_dim)``, no bias
o_proj : nn.Linear ``(hidden_dim, hidden_dim)``, no bias
rope : RotaryPositionEmbedding
Shape
-----
Input : ``(batch, seq_len, hidden_dim)``
Output : ``(batch, seq_len, hidden_dim)``
"""
def __init__(
self,
config: CyberSLMConfig,
rope: Optional[RotaryPositionEmbedding] = None,
) -> None:
super().__init__()
self.hidden_dim = config.hidden_dim
self.num_heads = config.num_heads
self.head_dim = config.head_dim
self.scale = 1.0 / math.sqrt(self.head_dim)
self.attn_dropout_p = config.attn_dropout
# ------------------------------------------------------------------ #
# Projection layers — no bias (modern practice, saves ~4×384 params) #
# ------------------------------------------------------------------ #
self.q_proj = nn.Linear(config.hidden_dim, config.hidden_dim, bias=False)
self.k_proj = nn.Linear(config.hidden_dim, config.hidden_dim, bias=False)
self.v_proj = nn.Linear(config.hidden_dim, config.hidden_dim, bias=False)
self.o_proj = nn.Linear(config.hidden_dim, config.hidden_dim, bias=False)
# ------------------------------------------------------------------ #
# RoPE cache #
# ------------------------------------------------------------------ #
# The cos/sin tables depend only on (head_dim, max_seq_len, base), so
# every layer's would be byte-identical. CyberSLM builds ONE and passes
# it in; previously each of the 12 layers constructed its own, costing
# ~12 MB of duplicated buffers. Falls back to building its own so the
# module stays usable standalone (tests, ablations).
self.rope = rope if rope is not None else RotaryPositionEmbedding(
head_dim=config.head_dim,
max_seq_len=config.max_seq_len,
base=config.rope_base,
)
# Causality is enforced by scaled_dot_product_attention(is_causal=...)
# rather than a materialised (max_seq_len × max_seq_len) mask buffer,
# which previously cost ~67 MB per layer.
def forward(
self,
x: Tensor,
attention_mask: Optional[Tensor] = None,
return_attn_weights: bool = False,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
use_cache: bool = False,
) -> Tuple[Tensor, Optional[Tensor], Optional[Tuple[Tensor, Tensor]]]:
"""
Compute multi-head self attention.
Parameters
----------
x : Tensor
Input of shape ``(batch, seq_len, hidden_dim)``.
attention_mask : Optional[Tensor]
Key-padding mask of shape ``(batch, seq_len)`` with 1 for real
tokens and 0 for padding. When provided, padded keys are excluded
from every query's attention (in addition to the causal mask).
``None`` means no padding (the common training case with packed
sequences).
return_attn_weights : bool
If True, also return the attention weight matrix for inspection.
This forces the slower explicit-softmax path; leave False for
training so the fused kernel is used.
Returns
-------
output : Tensor
Shape ``(batch, seq_len, hidden_dim)``.
attn_weights : Optional[Tensor]
Shape ``(batch, num_heads, seq_len, seq_len)`` if
``return_attn_weights=True``, else ``None``.
present : Optional[Tuple[Tensor, Tensor]]
The concatenated ``(k, v)`` for this layer when ``use_cache=True``,
to be fed back on the next decoding step. ``None`` otherwise.
"""
B, T, _ = x.shape
# Number of tokens already in the cache == absolute position of x[0].
past_len = kv_cache[0].size(2) if kv_cache is not None else 0
# ------------------------------------------------------------------ #
# 1. Linear projections #
# ------------------------------------------------------------------ #
q = self.q_proj(x) # (B, T, hidden_dim)
k = self.k_proj(x)
v = self.v_proj(x)
# ------------------------------------------------------------------ #
# 2. Reshape to (B, H, T, head_dim) for multi-head computation #
# ------------------------------------------------------------------ #
q = q.view(B, T, self.num_heads, self.head_dim).transpose(1, 2) # (B,H,T,D)
k = k.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
# ------------------------------------------------------------------ #
# 3. Apply Rotary Position Embeddings to Q and K #
# ------------------------------------------------------------------ #
# offset=past_len so a cached decode step rotates the new token by its
# TRUE absolute position rather than position 0.
q, k = apply_rope(q, k, self.rope, offset=past_len)
# ------------------------------------------------------------------ #
# 3b. Prepend the cache. RoPE is applied to the new k BEFORE the
# concat, and cached keys were already rotated when they were first
# computed -- so each key keeps the rotation for its own position.
# ------------------------------------------------------------------ #
if kv_cache is not None:
k = torch.cat([kv_cache[0], k], dim=2)
v = torch.cat([kv_cache[1], v], dim=2)
present = (k, v) if use_cache else None
S = k.size(2) # total key length (past + current)
# ------------------------------------------------------------------ #
# 4. Build the additive key-padding bias (if any). #
# Shape broadcasts over heads and query positions: (B, 1, 1, T). #
# ------------------------------------------------------------------ #
pad_bias: Optional[Tensor] = None
if attention_mask is not None:
# 0 where padding → -inf added to those key columns.
pad = (attention_mask == 0)[:, None, None, :] # (B,1,1,S) bool
pad_bias = torch.zeros(
(B, 1, 1, pad.size(-1)), dtype=q.dtype, device=q.device
).masked_fill(pad, float("-inf"))
if not return_attn_weights:
# Fused, memory-efficient path (FlashAttention when available).
# is_causal=True applies the causal mask without materialising it.
if pad_bias is None and past_len == 0:
context = F.scaled_dot_product_attention(
q, k, v,
is_causal=True,
dropout_p=self.attn_dropout_p if self.training else 0.0,
)
elif pad_bias is None and T == 1:
# Single-token decode: every cached key is in the past, so the
# causal constraint is already satisfied and no mask is needed.
context = F.scaled_dot_product_attention(
q, k, v, dropout_p=0.0,
)
else:
# Combine causal + padding into one additive float mask.
# Query i sits at absolute position past_len + i and may attend
# to keys 0..past_len+i, so the triangle is offset by past_len.
causal = _causal_bias(T, S, past_len, q.dtype, q.device)
attn_bias = causal[None, None, :, :]
if pad_bias is not None:
attn_bias = attn_bias + pad_bias # (B,1,T,S)
context = F.scaled_dot_product_attention(
q, k, v,
attn_mask=attn_bias,
dropout_p=self.attn_dropout_p if self.training else 0.0,
)
attn_weights = None
else:
# Explicit path — needed only when the caller wants the weights.
scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale # (B,H,T,S)
causal = _causal_bias(T, S, past_len, scores.dtype, scores.device)
scores = scores + causal[None, None, :, :]
if pad_bias is not None:
scores = scores + pad_bias
attn_weights = F.softmax(scores, dim=-1, dtype=torch.float32)
if self.attn_dropout_p > 0.0 and self.training:
attn_weights = F.dropout(attn_weights, p=self.attn_dropout_p)
attn_weights = attn_weights.to(v.dtype)
context = torch.matmul(attn_weights, v) # (B,H,T,head_dim)
# ------------------------------------------------------------------ #
# 8. Merge heads: (B, H, T, D) → (B, T, H*D) = (B, T, hidden_dim) #
# ------------------------------------------------------------------ #
context = context.transpose(1, 2).contiguous().view(B, T, self.hidden_dim)
# ------------------------------------------------------------------ #
# 9. Output projection #
# ------------------------------------------------------------------ #
output = self.o_proj(context)
return output, attn_weights, present
def extra_repr(self) -> str:
return (
f"hidden_dim={self.hidden_dim}, "
f"num_heads={self.num_heads}, "
f"head_dim={self.head_dim}"
)

144
cyberslm/model/block.py Normal file
View File

@@ -0,0 +1,144 @@
"""
Transformer Decoder Block
==========================
A single Pre-Norm residual decoder block consisting of:
1. RMSNorm → Multi-Head Self Attention → residual add
2. RMSNorm → SwiGLU FFN → residual add
Pre-Norm architecture
---------------------
Post-Norm (original Transformer):
x = LayerNorm(x + SubLayer(x))
Pre-Norm (modern: GPT-2 onward, LLaMA, etc.):
x = x + SubLayer(LayerNorm(x))
Pre-Norm is strongly preferred for deep networks because:
- Gradients flow through the residual connection bypassing the
normalisation, preventing vanishing gradients in very deep stacks.
- Training is more stable without learning-rate warmup tricks.
- Final RMSNorm on the output is added at the model level (not here)
to normalise the final residual stream before the LM head.
Residual stream
---------------
The residual stream x ∈ ^{B×T×d} is the backbone of the model. Each
sub-layer reads from it, computes a delta, and adds back:
Δ_attn = MHSA( RMSNorm(x) )
x = x + Δ_attn
Δ_ffn = FFN( RMSNorm(x) )
x = x + Δ_ffn
This additive structure means the gradient of the loss with respect to
early layers contains a direct path through the identity (residual),
enabling reliable training of 12+ layer models.
"""
from __future__ import annotations
from typing import Optional, Tuple
import torch
import torch.nn as nn
from torch import Tensor
from cyberslm.model.config import CyberSLMConfig
from cyberslm.model.norm import RMSNorm
from cyberslm.model.attention import MultiHeadSelfAttention
from cyberslm.model.ffn import SwiGLUFeedForward
class DecoderBlock(nn.Module):
"""
Pre-Norm Transformer Decoder Block.
Parameters
----------
config : CyberSLMConfig
Validated model configuration.
layer_idx : int
Zero-based index of this block in the stack (used for display only).
Sub-modules
-----------
attn_norm : RMSNorm
Normalises the residual stream before attention.
attn : MultiHeadSelfAttention
Self attention with RoPE and causal masking.
ffn_norm : RMSNorm
Normalises the residual stream before the FFN.
ffn : SwiGLUFeedForward
SwiGLU position-wise feed-forward network.
Shape
-----
Input : ``(batch, seq_len, hidden_dim)``
Output : ``(batch, seq_len, hidden_dim)``
"""
def __init__(
self,
config: CyberSLMConfig,
layer_idx: int = 0,
rope=None,
) -> None:
super().__init__()
self.layer_idx = layer_idx
# Pre-norm before attention.
self.attn_norm = RMSNorm(config.hidden_dim, eps=config.norm_eps)
# Multi-head self attention (owns RoPE + causal mask buffers).
self.attn = MultiHeadSelfAttention(config, rope=rope)
# Pre-norm before FFN.
self.ffn_norm = RMSNorm(config.hidden_dim, eps=config.norm_eps)
# SwiGLU feed-forward.
self.ffn = SwiGLUFeedForward(config)
def forward(
self,
x: Tensor,
attention_mask: Optional[Tensor] = None,
return_attn_weights: bool = False,
kv_cache: Optional[Tuple[Tensor, Tensor]] = None,
use_cache: bool = False,
) -> Tuple[Tensor, Optional[Tensor], Optional[Tuple[Tensor, Tensor]]]:
"""
Apply one Pre-Norm residual decoder block.
Parameters
----------
x : Tensor
Residual stream of shape ``(batch, seq_len, hidden_dim)``.
attention_mask : Optional[Tensor]
Key-padding mask ``(batch, seq_len)`` (1=keep, 0=pad), forwarded
to the attention sub-layer. ``None`` for packed/unpadded batches.
return_attn_weights : bool
Propagated to the attention sub-layer.
Returns
-------
x : Tensor
Updated residual stream, same shape as input.
attn_weights : Optional[Tensor]
Attention weights if requested, else None.
"""
# ---- Attention sub-layer ----------------------------------------- #
attn_out, attn_weights, present = self.attn(
self.attn_norm(x),
attention_mask=attention_mask,
return_attn_weights=return_attn_weights,
kv_cache=kv_cache,
use_cache=use_cache,
)
x = x + attn_out
# ---- FFN sub-layer ------------------------------------------------ #
x = x + self.ffn(self.ffn_norm(x))
return x, attn_weights, present
def extra_repr(self) -> str:
return f"layer_idx={self.layer_idx}"

234
cyberslm/model/config.py Normal file
View File

@@ -0,0 +1,234 @@
"""
CyberSLM Model Configuration
=============================
Defines the complete, validated configuration for the CyberSLM decoder-only
transformer. All hyperparameters are frozen at construction time and validated
for mathematical consistency before any model component is instantiated.
Architecture summary
--------------------
- Hidden dim : 384
- Decoder layers : 12
- Attention heads : 6
- Head dim : 64 (hidden_dim / num_heads = 384 / 6 = 64)
- FFN inner dim : 1024 (SwiGLU gate + value = 2 × 1024 → projects back to 384)
- Vocab size : 32 000
- Max context : 4 096
- Approx params : 33.53 M
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Optional
@dataclass(frozen=True)
class CyberSLMConfig:
"""
Immutable configuration for the CyberSLM decoder-only transformer.
All fields are set once at construction; the frozen dataclass guarantees
no accidental mutation during training. Call :meth:`validate` immediately
after construction or use the convenience constructor
:func:`default_config`.
Attributes
----------
vocab_size : int
Number of tokens in the SentencePiece BPE vocabulary.
max_seq_len : int
Maximum token sequence length (context window).
hidden_dim : int
Embedding dimension ``d_model``.
num_layers : int
Number of stacked transformer decoder blocks.
num_heads : int
Number of attention heads. Must evenly divide ``hidden_dim``.
head_dim : int
Dimension of each attention head. Must equal ``hidden_dim // num_heads``.
ffn_hidden_dim : int
Inner dimension of the SwiGLU feed-forward network.
The gate and value projections each map hidden_dim → ffn_hidden_dim,
and the output projection maps ffn_hidden_dim → hidden_dim.
rope_base : int
Base for Rotary Position Embedding frequency computation (θ = 10 000).
norm_eps : float
Epsilon added inside RMSNorm to prevent division by zero.
tie_weights : bool
When True the output projection shares weights with the token embedding.
bias : bool
When True linear layers include a bias term (False = modern practice).
dropout : float
Residual / feed-forward dropout probability (0.0 = disabled).
attn_dropout : float
Attention weight dropout probability (0.0 = disabled).
pad_token_id : Optional[int]
Token ID used for padding; None if the dataset never pads.
bos_token_id : Optional[int]
Beginning-of-sequence token ID.
eos_token_id : Optional[int]
End-of-sequence token ID.
"""
# ------------------------------------------------------------------ #
# Vocabulary & sequence #
# ------------------------------------------------------------------ #
vocab_size: int = 32_000
max_seq_len: int = 4_096
# ------------------------------------------------------------------ #
# Transformer dimensions #
# ------------------------------------------------------------------ #
hidden_dim: int = 384
num_layers: int = 12
num_heads: int = 6
head_dim: int = 64 # must equal hidden_dim // num_heads
ffn_hidden_dim: int = 1_024
# ------------------------------------------------------------------ #
# Positional encoding #
# ------------------------------------------------------------------ #
rope_base: int = 10_000
# ------------------------------------------------------------------ #
# Normalization #
# ------------------------------------------------------------------ #
norm_eps: float = 1e-6
# ------------------------------------------------------------------ #
# Architecture flags #
# ------------------------------------------------------------------ #
tie_weights: bool = True
bias: bool = False
dropout: float = 0.0
attn_dropout: float = 0.0
# ------------------------------------------------------------------ #
# Special token IDs (set by tokenizer integration layer) #
# ------------------------------------------------------------------ #
pad_token_id: Optional[int] = None
bos_token_id: Optional[int] = 2 # real SentencePiece BOS id
eos_token_id: Optional[int] = 3 # real SentencePiece EOS id
# ------------------------------------------------------------------ #
# Validation #
# ------------------------------------------------------------------ #
def validate(self) -> "CyberSLMConfig":
"""
Assert mathematical consistency of every hyperparameter.
Raises
------
ValueError
If any hyperparameter violates an architectural constraint.
Returns
-------
CyberSLMConfig
Self, to allow chaining: ``cfg = CyberSLMConfig().validate()``.
"""
errors: list[str] = []
# Positivity checks
for name, value in [
("vocab_size", self.vocab_size),
("max_seq_len", self.max_seq_len),
("hidden_dim", self.hidden_dim),
("num_layers", self.num_layers),
("num_heads", self.num_heads),
("head_dim", self.head_dim),
("ffn_hidden_dim", self.ffn_hidden_dim),
("rope_base", self.rope_base),
]:
if value <= 0:
errors.append(f"{name} must be positive, got {value}")
# Attention head consistency
if self.hidden_dim % self.num_heads != 0:
errors.append(
f"hidden_dim ({self.hidden_dim}) must be divisible by "
f"num_heads ({self.num_heads})"
)
expected_head_dim = self.hidden_dim // self.num_heads
if self.head_dim != expected_head_dim:
errors.append(
f"head_dim ({self.head_dim}) must equal "
f"hidden_dim // num_heads = {expected_head_dim}"
)
# RoPE requires even head_dim (pairs of sin/cos)
if self.head_dim % 2 != 0:
errors.append(
f"head_dim ({self.head_dim}) must be even for RoPE"
)
# Dropout bounds
for name, value in [("dropout", self.dropout), ("attn_dropout", self.attn_dropout)]:
if not (0.0 <= value < 1.0):
errors.append(f"{name} must be in [0, 1), got {value}")
# norm_eps positivity
if self.norm_eps <= 0.0:
errors.append(f"norm_eps must be positive, got {self.norm_eps}")
if errors:
raise ValueError(
"CyberSLMConfig validation failed:\n"
+ "\n".join(f"{e}" for e in errors)
)
return self
# ------------------------------------------------------------------ #
# Derived properties #
# ------------------------------------------------------------------ #
@property
def total_attention_dim(self) -> int:
"""``num_heads × head_dim`` — equals ``hidden_dim`` by construction."""
return self.num_heads * self.head_dim
@property
def rope_half_dim(self) -> int:
"""Number of frequency pairs in RoPE (``head_dim // 2``)."""
return self.head_dim // 2
# ------------------------------------------------------------------ #
# Display #
# ------------------------------------------------------------------ #
def __str__(self) -> str:
lines = [
"CyberSLMConfig",
"=" * 40,
f" vocab_size : {self.vocab_size:,}",
f" max_seq_len : {self.max_seq_len:,}",
f" hidden_dim : {self.hidden_dim}",
f" num_layers : {self.num_layers}",
f" num_heads : {self.num_heads}",
f" head_dim : {self.head_dim}",
f" ffn_hidden_dim : {self.ffn_hidden_dim}",
f" rope_base : {self.rope_base}",
f" norm_eps : {self.norm_eps}",
f" tie_weights : {self.tie_weights}",
f" bias : {self.bias}",
f" dropout : {self.dropout}",
f" attn_dropout : {self.attn_dropout}",
]
return "\n".join(lines)
def default_config() -> CyberSLMConfig:
"""
Return the validated default CyberSLM configuration.
This is the single source of truth for all training runs.
All hyperparameters match the finalized architecture specification.
Returns
-------
CyberSLMConfig
A validated, immutable configuration object.
"""
cfg = CyberSLMConfig()
cfg.validate()
return cfg

127
cyberslm/model/ffn.py Normal file
View File

@@ -0,0 +1,127 @@
"""
SwiGLU Feed-Forward Network (FFN)
===================================
Reference: "GLU Variants Improve Transformer" (Noam Shazeer, 2020)
https://arxiv.org/abs/2002.05202
Mathematical definition
-----------------------
Standard FFN (for contrast):
FFN(x) = activation(x W₁) W₂
SwiGLU FFN:
SwiGLU(x) = (x W_gate ⊙ swish(x W_gate)) W₂ ← WRONG shorthand
Correct form with separate gate and value projections:
gate(x) = x W_gate ∈ ^{B×T×ffn_dim}
val(x) = x W_val ∈ ^{B×T×ffn_dim}
hidden = swish(gate(x)) ⊙ val(x)
out = hidden W_out ∈ ^{B×T×hidden_dim}
where swish(z) = z · sigmoid(z) = z · σ(z).
Why SwiGLU
----------
- The gating mechanism (⊙) gives the network a multiplicative path to
control information flow — intuitively "how much" of each feature passes
through at each position.
- swish is smooth and non-monotonic, empirically outperforming ReLU and
GELU in large-scale experiments (PaLM, LLaMA, etc.).
- The factored W_gate / W_val structure adds one extra matrix but improves
quality relative to a standard 2-layer FFN of the same parameter budget.
Parameter count
---------------
Three matrices: W_gate, W_val, W_out
W_gate : hidden_dim × ffn_hidden_dim (384 × 1024 = 393 216)
W_val : hidden_dim × ffn_hidden_dim (384 × 1024 = 393 216)
W_out : ffn_hidden_dim × hidden_dim (1024 × 384 = 393 216)
Total per layer: 1 179 648 ≈ 1.18 M
No bias on any projection (consistent with modern practice).
Numerical stability
-------------------
- swish(z) = z · σ(z) is numerically safe for all real z.
- σ(z) = 1/(1+exp(z)) in PyTorch uses a numerically stable implementation.
- The element-wise product of two bounded (after sigmoid) quantities does
not amplify values explosively.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from cyberslm.model.config import CyberSLMConfig
class SwiGLUFeedForward(nn.Module):
"""
SwiGLU Feed-Forward Network.
Consists of three bias-free linear projections:
``gate_proj``, ``val_proj``, and ``out_proj``.
Parameters
----------
config : CyberSLMConfig
Validated model configuration.
Shape
-----
Input : ``(batch, seq_len, hidden_dim)``
Output : ``(batch, seq_len, hidden_dim)``
"""
def __init__(self, config: CyberSLMConfig) -> None:
super().__init__()
self.hidden_dim = config.hidden_dim
self.ffn_hidden_dim = config.ffn_hidden_dim
# Gate projection: produces the gating signal fed through swish.
self.gate_proj = nn.Linear(
config.hidden_dim, config.ffn_hidden_dim, bias=False
)
# Value projection: produces the value signal gated element-wise.
self.val_proj = nn.Linear(
config.hidden_dim, config.ffn_hidden_dim, bias=False
)
# Output projection: maps back to residual stream dimension.
self.out_proj = nn.Linear(
config.ffn_hidden_dim, config.hidden_dim, bias=False
)
def forward(self, x: Tensor) -> Tensor:
"""
Apply SwiGLU feed-forward transformation.
Parameters
----------
x : Tensor
Input of shape ``(batch, seq_len, hidden_dim)``.
Returns
-------
Tensor
Output of shape ``(batch, seq_len, hidden_dim)``.
"""
# gate: (B, T, ffn_hidden_dim)
# val: (B, T, ffn_hidden_dim)
gate = self.gate_proj(x)
val = self.val_proj(x)
# SwiGLU: swish(gate) ⊙ val
# F.silu is swish: silu(z) = z * sigmoid(z)
hidden = F.silu(gate) * val # (B, T, ffn_hidden_dim)
# Project back to hidden_dim
return self.out_proj(hidden) # (B, T, hidden_dim)
def extra_repr(self) -> str:
return (
f"hidden_dim={self.hidden_dim}, "
f"ffn_hidden_dim={self.ffn_hidden_dim}"
)

514
cyberslm/model/model.py Normal file
View File

@@ -0,0 +1,514 @@
"""
CyberSLM — Complete Decoder-Only Language Model
================================================
Assembles all components into the full model:
Token Embedding → Decoder Stack (×12) → Final RMSNorm → LM Head
Weight tying
------------
The LM head (output projection that maps hidden_dim → vocab_size) shares its
weight matrix with the token embedding (vocab_size × hidden_dim).
Mathematical justification: both the embedding matrix E and the unembedding
matrix U operate in the same semantic space. Tying U = Eᵀ forces consistency
("a token's output representation should be similar to its input representation"),
reduces parameters by vocab_size × hidden_dim = 32 000 × 384 ≈ 12.3 M, and
empirically improves perplexity on small models.
Parameter count breakdown
--------------------------
Component Params
---------------------------------------- ------
Token embedding (vocab × hidden) 12 288 000
↳ shared with LM head (no extra cost) 0
Decoder blocks × 12:
attn_norm (RMSNorm) 384 ×12 = 4 608
attn Q/K/V/O proj 589 824 ×12 = 7 077 888
ffn_norm (RMSNorm) 384 ×12 = 4 608
ffn gate/val/out 1 179 648 ×12 = 14 155 776
Final RMSNorm 384
---------------------------------------- ------
Total ≈ 33 531 264 (≈33.53 M)
Note: RoPE buffers and causal mask buffers are NOT parameters.
Initialisation
--------------
- Embeddings: N(0, 0.02) — small but non-zero, standard practice.
- All linear weights: N(0, 0.02)
- All RMSNorm γ: 1.0 (already set by RMSNorm.__init__)
- Output projection weights = Embedding weights (weight tying).
- Scaled output projections: attention o_proj and FFN out_proj are
scaled by 1/√(2·num_layers) to prevent residual stream variance
from growing with depth (following GPT-2 / LLaMA init practice).
"""
from __future__ import annotations
import math
from typing import Dict, List, Optional, Tuple
import torch
import torch.nn as nn
from torch import Tensor
from cyberslm.model.config import CyberSLMConfig, default_config
from cyberslm.model.block import DecoderBlock
from cyberslm.model.norm import RMSNorm
from cyberslm.model.rope import RotaryPositionEmbedding
class CyberSLM(nn.Module):
"""
CyberSLM Decoder-Only Transformer.
Parameters
----------
config : CyberSLMConfig
Validated model configuration.
Attributes
----------
config : CyberSLMConfig
embedding : nn.Embedding
Token embedding table, shape ``(vocab_size, hidden_dim)``.
layers : nn.ModuleList[DecoderBlock]
Stack of ``num_layers`` decoder blocks.
final_norm : RMSNorm
Applied to the residual stream after the last block.
lm_head : nn.Linear
Projects hidden_dim → vocab_size. Weight tied to ``embedding``.
"""
def __init__(self, config: CyberSLMConfig) -> None:
super().__init__()
config.validate()
self.config = config
# ------------------------------------------------------------------ #
# Token embedding #
# ------------------------------------------------------------------ #
self.embedding = nn.Embedding(config.vocab_size, config.hidden_dim)
# ------------------------------------------------------------------ #
# Decoder stack #
# ------------------------------------------------------------------ #
# One RoPE table shared by every layer (identical by construction).
self.rope = RotaryPositionEmbedding(
head_dim=config.head_dim,
max_seq_len=config.max_seq_len,
base=config.rope_base,
)
self.layers = nn.ModuleList(
[
DecoderBlock(config, layer_idx=i, rope=self.rope)
for i in range(config.num_layers)
]
)
# ------------------------------------------------------------------ #
# Final normalisation #
# ------------------------------------------------------------------ #
self.final_norm = RMSNorm(config.hidden_dim, eps=config.norm_eps)
# ------------------------------------------------------------------ #
# Language model head #
# ------------------------------------------------------------------ #
# bias=False: unembedding never needs a bias term.
self.lm_head = nn.Linear(config.hidden_dim, config.vocab_size, bias=False)
# ------------------------------------------------------------------ #
# Weight tying: lm_head.weight ≡ embedding.weight #
# ------------------------------------------------------------------ #
if config.tie_weights:
self.lm_head.weight = self.embedding.weight
# ------------------------------------------------------------------ #
# Parameter initialisation #
# ------------------------------------------------------------------ #
self._init_weights()
# ---------------------------------------------------------------------- #
# Initialisation #
# ---------------------------------------------------------------------- #
def _init_weights(self) -> None:
"""
Initialise all parameters with production-quality values.
Strategy
--------
- Embedding : N(0, 0.02)
- All Linear weights : N(0, 0.02)
- o_proj and out_proj: scaled down by 1/√(2·L) where L = num_layers
to stabilise the residual stream variance at initialisation.
- All RMSNorm γ : 1.0 (already set in RMSNorm.__init__)
- All biases : 0.0 (none exist in this config)
"""
std = 0.02
scaled_std = std / math.sqrt(2.0 * self.config.num_layers)
for name, module in self.named_modules():
if isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=std)
elif isinstance(module, nn.Linear):
# Scaled init for residual output projections.
if name.endswith("o_proj") or name.endswith("out_proj"):
nn.init.normal_(module.weight, mean=0.0, std=scaled_std)
else:
nn.init.normal_(module.weight, mean=0.0, std=std)
if module.bias is not None:
nn.init.zeros_(module.bias)
# Weight tying must be re-applied after init because _init_weights
# initialised embedding.weight; lm_head.weight already points to the
# same tensor (Python object reference), so no extra step needed.
# Verify it is still tied.
if self.config.tie_weights:
assert self.lm_head.weight is self.embedding.weight, (
"Weight tying broken after _init_weights"
)
# ---------------------------------------------------------------------- #
# Forward pass #
# ---------------------------------------------------------------------- #
def forward(
self,
input_ids: Tensor,
attention_mask: Optional[Tensor] = None,
return_all_attn_weights: bool = False,
) -> Tuple[Tensor, List[Optional[Tensor]]]:
"""
Run the full forward pass.
Parameters
----------
input_ids : Tensor
Long tensor of shape ``(batch, seq_len)`` with token IDs in
``[0, vocab_size)``.
attention_mask : Optional[Tensor]
Optional key-padding mask ``(batch, seq_len)`` (1=keep, 0=pad).
Pass this when batches contain right-padded sequences so padded
positions do not corrupt real tokens; leave ``None`` for packed,
unpadded training batches.
return_all_attn_weights : bool
If True, collect and return attention weights from every layer.
Disabled by default for training efficiency.
Returns
-------
logits : Tensor
Shape ``(batch, seq_len, vocab_size)``. Raw (pre-softmax) scores.
all_attn_weights : List[Optional[Tensor]]
One entry per decoder block; each is either the attention weight
tensor or ``None``.
Notes
-----
For language model training the standard loss is:
loss = cross_entropy(logits[:, :-1].reshape(-1, V),
input_ids[:, 1:].reshape(-1))
where we predict the next token at every position.
"""
# ------------------------------------------------------------------ #
# 1. Token embedding #
# ------------------------------------------------------------------ #
x = self.embedding(input_ids) # (B, T, hidden_dim)
# ------------------------------------------------------------------ #
# 2. Decoder stack #
# ------------------------------------------------------------------ #
all_attn_weights: List[Optional[Tensor]] = []
for block in self.layers:
# use_cache=False -> `present` is None; forward() deliberately keeps
# its (logits, attn_weights) return signature unchanged. Cached
# decoding lives in generate() instead of overloading this method.
x, attn_w, _ = block(
x,
attention_mask=attention_mask,
return_attn_weights=return_all_attn_weights,
)
all_attn_weights.append(attn_w)
# ------------------------------------------------------------------ #
# 3. Final normalisation #
# ------------------------------------------------------------------ #
x = self.final_norm(x) # (B, T, hidden_dim)
# ------------------------------------------------------------------ #
# 4. LM head (weight-tied unembedding) #
# ------------------------------------------------------------------ #
logits = self.lm_head(x) # (B, T, vocab_size)
return logits, all_attn_weights
# ---------------------------------------------------------------------- #
# Cached autoregressive generation #
# ---------------------------------------------------------------------- #
@torch.no_grad()
def generate(
self,
input_ids: Tensor,
max_new_tokens: int = 256,
temperature: float = 0.0,
top_k: int = 0,
top_p: float = 1.0,
repetition_penalty: float = 1.0,
eos_id: Optional[int] = None,
) -> Tensor:
"""
Generate continuations using a KV cache.
Why this exists
---------------
The previous generation loop re-ran the whole 12-layer stack over the
entire prefix for every single token, making decoding O(n^2) in
sequence length. With a cache each step attends over the cached keys and
only computes the new token, which is O(n) overall.
Sampling is applied per row, so batched prompts are supported. Rows that
have emitted ``eos_id`` are frozen (further tokens are forced to
``eos_id``) and generation stops once every row is finished.
Parameters
----------
input_ids : Tensor ``(batch, prompt_len)`` of token ids.
temperature : 0.0 selects greedy argmax; >0 samples.
top_k / top_p : 0 and 1.0 respectively disable the filter.
repetition_penalty : >1.0 divides logits of already-present tokens.
eos_id : stop token; ``None`` means never stop early.
Returns
-------
Tensor ``(batch, prompt_len + generated)`` including the prompt.
"""
self.eval()
device = input_ids.device
B = input_ids.size(0)
max_ctx = self.config.max_seq_len
if input_ids.size(1) >= max_ctx:
input_ids = input_ids[:, -(max_ctx - 1):]
caches: List[Optional[tuple]] = [None] * len(self.layers)
finished = torch.zeros(B, dtype=torch.bool, device=device)
out = input_ids
cur = input_ids
for _ in range(max_new_tokens):
if out.size(1) >= max_ctx:
break
h = self.embedding(cur)
new_caches = []
for block, layer_cache in zip(self.layers, caches):
h, _, present = block(h, kv_cache=layer_cache, use_cache=True)
new_caches.append(present)
caches = new_caches
logits = self.lm_head(self.final_norm(h))[:, -1, :].float()
if repetition_penalty != 1.0:
for b in range(B):
seen = torch.unique(out[b])
lg = logits[b, seen]
# Divide positives, multiply negatives, so the penalty always
# pushes a token DOWN regardless of its logit's sign.
logits[b, seen] = torch.where(
lg > 0, lg / repetition_penalty, lg * repetition_penalty
)
if temperature == 0.0:
nxt = logits.argmax(dim=-1)
else:
logits = logits / temperature
if top_k > 0:
k = min(top_k, logits.size(-1))
thresh = torch.topk(logits, k, dim=-1).values[:, -1, None]
logits = logits.masked_fill(logits < thresh, float("-inf"))
if top_p < 1.0:
srt, idx = torch.sort(logits, descending=True, dim=-1)
probs = torch.softmax(srt, dim=-1)
cum = probs.cumsum(dim=-1) - probs # prob mass strictly before this token
srt = srt.masked_fill(cum > top_p, float("-inf"))
logits = torch.full_like(logits, float("-inf")).scatter(1, idx, srt)
nxt = torch.multinomial(torch.softmax(logits, dim=-1), 1).squeeze(-1)
if eos_id is not None:
nxt = torch.where(finished, torch.full_like(nxt, eos_id), nxt)
finished = finished | (nxt == eos_id)
cur = nxt.unsqueeze(1)
out = torch.cat([out, cur], dim=1)
if eos_id is not None and bool(finished.all()):
break
return out
# ---------------------------------------------------------------------- #
# Convenience: next-token logits #
# ---------------------------------------------------------------------- #
def get_next_token_logits(self, input_ids: Tensor) -> Tensor:
"""
Return logits for the next token after the last input position.
Parameters
----------
input_ids : Tensor
Shape ``(batch, seq_len)``.
Returns
-------
Tensor
Shape ``(batch, vocab_size)``.
"""
logits, _ = self.forward(input_ids)
return logits[:, -1, :] # (B, vocab_size)
# --------------------------------------------------------------------------- #
# Parameter counting #
# --------------------------------------------------------------------------- #
def count_parameters(model: nn.Module) -> Dict[str, int]:
"""
Count trainable and total parameters.
Because of weight tying, lm_head.weight is counted only once
(it shares storage with embedding.weight).
Parameters
----------
model : nn.Module
Returns
-------
dict with keys:
``total`` — total parameter elements (no double-counting)
``trainable`` — trainable parameter elements
"""
seen: set = set()
total = 0
trainable = 0
for param in model.parameters():
# data_ptr() is unique per underlying storage tensor.
if param.data_ptr() in seen:
continue
seen.add(param.data_ptr())
n = param.numel()
total += n
if param.requires_grad:
trainable += n
return {"total": total, "trainable": trainable}
# --------------------------------------------------------------------------- #
# Model summary #
# --------------------------------------------------------------------------- #
def model_summary(model: CyberSLM) -> str:
"""
Return a human-readable model summary string.
Parameters
----------
model : CyberSLM
Returns
-------
str
Formatted multi-line summary including per-component parameter counts.
"""
cfg = model.config
param_info = count_parameters(model)
lines = [
"=" * 60,
f" CyberSLM Model Summary",
"=" * 60,
f" Architecture : Decoder-only Transformer",
f" Hidden dim : {cfg.hidden_dim}",
f" Num layers : {cfg.num_layers}",
f" Num heads : {cfg.num_heads}",
f" Head dim : {cfg.head_dim}",
f" FFN hidden dim : {cfg.ffn_hidden_dim}",
f" Vocab size : {cfg.vocab_size:,}",
f" Max seq len : {cfg.max_seq_len:,}",
f" Weight tied : {cfg.tie_weights}",
f" RoPE base : {cfg.rope_base}",
"-" * 60,
f" Total params : {param_info['total']:>14,}",
f" Trainable : {param_info['trainable']:>14,}",
"-" * 60,
" Per-component:",
]
# Embedding
emb_p = model.embedding.weight.numel()
lines.append(f" Embedding : {emb_p:>12,}")
# Per-block breakdown (just the first block, all identical)
block = model.layers[0]
attn_norm_p = sum(p.numel() for p in block.attn_norm.parameters())
attn_p = sum(p.numel() for p in block.attn.parameters())
ffn_norm_p = sum(p.numel() for p in block.ffn_norm.parameters())
ffn_p = sum(p.numel() for p in block.ffn.parameters())
block_total = attn_norm_p + attn_p + ffn_norm_p + ffn_p
lines.append(f" Decoder block (×{cfg.num_layers:2d}) : {block_total:>12,} per block")
lines.append(f" attn_norm : {attn_norm_p:>12,}")
lines.append(f" attention (Q/K/V/O) : {attn_p:>12,}")
lines.append(f" ffn_norm : {ffn_norm_p:>12,}")
lines.append(f" ffn (gate/val/out) : {ffn_p:>12,}")
lines.append(f" Decoder stack total : {block_total * cfg.num_layers:>12,}")
# Final norm
final_norm_p = sum(p.numel() for p in model.final_norm.parameters())
lines.append(f" Final RMSNorm : {final_norm_p:>12,}")
# LM head — note: weight tied, so 0 additional params
lm_tied_note = " (weight-tied, no extra params)" if cfg.tie_weights else ""
lines.append(f" LM head : {0:>12,}{lm_tied_note}")
lines.append("=" * 60)
return "\n".join(lines)
# --------------------------------------------------------------------------- #
# Model builder #
# --------------------------------------------------------------------------- #
def build_model(
config: Optional[CyberSLMConfig] = None,
device: Optional[torch.device] = None,
) -> CyberSLM:
"""
Build, initialise, and optionally place the CyberSLM model.
Parameters
----------
config : CyberSLMConfig, optional
Validated config. Uses :func:`default_config` if None.
device : torch.device, optional
Target device. Stays on CPU if None.
Returns
-------
CyberSLM
Fully initialised model ready for training.
"""
if config is None:
config = default_config()
else:
config.validate()
model = CyberSLM(config)
if device is not None:
model = model.to(device)
return model

114
cyberslm/model/norm.py Normal file
View File

@@ -0,0 +1,114 @@
"""
RMSNorm — Root Mean Square Layer Normalisation
===============================================
Reference: "Root Mean Square Layer Normalization" (Zhang & Sennrich, 2019)
https://arxiv.org/abs/1910.07467
Mathematical definition
-----------------------
Given an input vector **x** ∈ ^d:
RMS(x) = sqrt( (1/d) * Σ xᵢ² + ε )
RMSNorm(x) = (x / RMS(x)) * γ
where γ^d is a learned per-channel scale (initialised to 1.0) and
ε > 0 is a small constant for numerical stability.
Key differences from LayerNorm
-------------------------------
- No mean subtraction (no re-centering step).
- No learned bias β (the bias-free variant).
- ~30 % fewer operations than LayerNorm, which matters across 12 layers.
- Empirically matches or exceeds LayerNorm in modern transformer training.
Numerical stability
-------------------
- The RMS is computed in float32 regardless of input dtype. This prevents
underflow/overflow when activations are in bfloat16 or float16. For our
FP32 training runs this cast is a no-op but it is correct and future-proof.
- ε = 1e-6 (default) prevents division by zero even for near-zero inputs.
- The scale γ is cast back to the input dtype before multiplication.
"""
from __future__ import annotations
import torch
import torch.nn as nn
from torch import Tensor
class RMSNorm(nn.Module):
"""
Root Mean Square Layer Normalisation without mean-centering or bias.
Parameters
----------
dim : int
Feature dimension to normalise over (last dimension of the input).
eps : float
Small constant added to the RMS denominator for numerical stability.
Defaults to 1e-6.
Shape
-----
Input : ``(*, dim)`` — any leading batch / sequence dimensions.
Output : ``(*, dim)`` — same shape as input.
Examples
--------
>>> norm = RMSNorm(384, eps=1e-6)
>>> x = torch.randn(2, 512, 384)
>>> y = norm(x)
>>> y.shape
torch.Size([2, 512, 384])
"""
def __init__(self, dim: int, eps: float = 1e-6) -> None:
super().__init__()
if dim <= 0:
raise ValueError(f"dim must be positive, got {dim}")
if eps <= 0.0:
raise ValueError(f"eps must be positive, got {eps}")
self.dim = dim
self.eps = eps
# Learned per-channel scale, initialised to 1 (identity transform).
self.weight = nn.Parameter(torch.ones(dim))
def _compute_rms(self, x: Tensor) -> Tensor:
"""
Compute the RMS over the last dimension.
Always promotes to float32 to prevent numerical issues with
reduced-precision dtypes. For FP32 training this is a no-op.
Returns
-------
Tensor
Shape ``(*, 1)`` — one RMS value per token position.
"""
return x.float().pow(2).mean(dim=-1, keepdim=True).add(self.eps).sqrt()
def forward(self, x: Tensor) -> Tensor:
"""
Normalise ``x`` by its per-token root mean square.
Parameters
----------
x : Tensor
Input of shape ``(*, dim)``.
Returns
-------
Tensor
Normalised output of the same shape and dtype as ``x``.
"""
rms = self._compute_rms(x)
# Normalise in float32, then cast back to original dtype.
x_normed = x.float() / rms
# Scale by learned weights (cast to input dtype for type safety).
return (x_normed * self.weight.float()).to(x.dtype)
def extra_repr(self) -> str:
return f"dim={self.dim}, eps={self.eps}"

291
cyberslm/model/rope.py Normal file
View File

@@ -0,0 +1,291 @@
"""
Rotary Position Embedding (RoPE)
=================================
Reference: "RoFormer: Enhanced Transformer with Rotary Position Embedding"
(Su et al., 2021) — https://arxiv.org/abs/2104.09864
Mathematical definition
-----------------------
For a query (or key) vector **q** at position ``m`` with head dimension ``d``:
1. Partition **q** into pairs: (q₁, q₂), (q₃, q₄), ..., (q_{d-1}, q_d).
2. For each pair index ``i ∈ {0, 1, ..., d/2 - 1}`` define the frequency:
θᵢ = 1 / base^(2i / d) (base = 10 000)
3. Apply a 2-D rotation to each pair at position ``m``:
R(m, θᵢ) · (q_{2i}, q_{2i+1}) =
(q_{2i} cos(m·θᵢ) q_{2i+1} sin(m·θᵢ),
q_{2i} sin(m·θᵢ) + q_{2i+1} cos(m·θᵢ))
This is equivalent to multiplying **q** (viewed as complex numbers) by
``exp(i · m · θ)``, which preserves the inner product of relative positions:
⟨R(m)q, R(n)k⟩ depends only on (m n),
giving translation-equivariant attention without absolute position tokens.
Efficient implementation
------------------------
The rotation can be expressed without complex arithmetic:
q_rot = [q_even · cos q_odd · sin,
q_even · sin + q_odd · cos]
where ``q_even = q[..., 0::2]``, ``q_odd = q[..., 1::2]``.
Numerically interleaved form (even/odd) vs. split-half form (first/second half)
are equivalent; we use the interleaved form for clarity.
Precomputation
--------------
``cos`` and ``sin`` tensors of shape ``(max_seq_len, head_dim // 2)`` are
computed once and registered as non-parameter buffers so they move with the
module (CPU ↔ GPU) and are not included in ``state_dict`` checkpoints.
Stability notes
---------------
- Frequencies are computed in float64 then cast to float32 to minimise
floating-point error in ``pow`` and ``arange``.
- All rotations are executed in float32 to prevent loss of precision.
"""
from __future__ import annotations
import math
from typing import Tuple
import torch
import torch.nn as nn
from torch import Tensor
class RotaryPositionEmbedding(nn.Module):
"""
Precomputed Rotary Position Embedding cache.
Registers ``cos`` and ``sin`` buffers of shape
``(max_seq_len, head_dim // 2)`` at construction time. Applying RoPE
to a query or key tensor costs only element-wise multiplications and
additions — no matrix multiplications.
Parameters
----------
head_dim : int
Dimension of each attention head. Must be even.
max_seq_len : int
Maximum sequence length to pre-compute. Sequences longer than this
will raise an error at runtime.
base : int
RoPE base frequency (10 000 in the original paper).
Shape of ``apply``
------------------
Input : ``(batch, num_heads, seq_len, head_dim)``
Output : ``(batch, num_heads, seq_len, head_dim)``
"""
def __init__(
self,
head_dim: int,
max_seq_len: int,
base: int = 10_000,
) -> None:
super().__init__()
if head_dim <= 0 or head_dim % 2 != 0:
raise ValueError(
f"head_dim must be a positive even integer, got {head_dim}"
)
if max_seq_len <= 0:
raise ValueError(f"max_seq_len must be positive, got {max_seq_len}")
if base <= 0:
raise ValueError(f"base must be positive, got {base}")
self.head_dim = head_dim
self.max_seq_len = max_seq_len
self.base = base
# Pre-compute and register buffers (not model parameters).
cos_cache, sin_cache = self._build_cache(head_dim, max_seq_len, base)
self.register_buffer("cos_cache", cos_cache, persistent=False)
self.register_buffer("sin_cache", sin_cache, persistent=False)
@staticmethod
def _build_cache(
head_dim: int,
max_seq_len: int,
base: int,
) -> Tuple[Tensor, Tensor]:
"""
Build ``(cos, sin)`` caches of shape ``(max_seq_len, head_dim // 2)``.
Computation is performed in float64 for precision, then cast to
float32 for storage.
Returns
-------
Tuple[Tensor, Tensor]
``cos_cache`` and ``sin_cache``, each of shape
``(max_seq_len, head_dim // 2)``.
"""
half_dim = head_dim // 2
# θᵢ = 1 / base^(2i / head_dim) for i ∈ {0, ..., half_dim - 1}
# Computed in float64 to avoid precision loss in the exponent.
inv_freq = 1.0 / (
base ** (torch.arange(0, head_dim, 2, dtype=torch.float64) / head_dim)
)
# Shape: (half_dim,)
# Position indices m ∈ {0, 1, ..., max_seq_len - 1}
positions = torch.arange(max_seq_len, dtype=torch.float64)
# Shape: (max_seq_len,)
# Outer product: angles[m, i] = m * θᵢ
angles = torch.outer(positions, inv_freq)
# Shape: (max_seq_len, half_dim)
cos_cache = angles.cos().to(torch.float32)
sin_cache = angles.sin().to(torch.float32)
return cos_cache, sin_cache
@staticmethod
def _rotate_half(x: Tensor) -> Tensor:
"""
Rotate the last dimension by interleaving even/odd pairs.
For input ``x`` of shape ``(..., head_dim)``:
x_even = x[..., 0::2] (positions 0, 2, 4, ...)
x_odd = x[..., 1::2] (positions 1, 3, 5, ...)
Returns ``[-x_odd, x_even]`` interleaved back into ``(..., head_dim)``.
This is the standard rotation that implements the complex-number trick.
Parameters
----------
x : Tensor
Shape ``(..., head_dim)`` where ``head_dim`` is even.
Returns
-------
Tensor
Same shape as ``x``.
"""
x_even = x[..., 0::2] # (..., head_dim // 2)
x_odd = x[..., 1::2] # (..., head_dim // 2)
# Interleave: stack along new dim then flatten.
rotated = torch.stack([-x_odd, x_even], dim=-1)
# Shape: (..., head_dim // 2, 2) → (..., head_dim)
return rotated.flatten(start_dim=-2)
def apply(self, x: Tensor, offset: int = 0) -> Tensor:
"""
Apply Rotary Position Embeddings to ``x``.
Parameters
----------
x : Tensor
Query or key tensor of shape
``(batch, num_heads, seq_len, head_dim)``.
offset : int
Absolute position of ``x[..., 0, :]`` in the full sequence.
This is what makes incremental decoding correct. With a KV cache the
model feeds one token at a time, so ``seq_len == 1`` and the naive
``cos_cache[:1]`` would rotate every generated token as if it were at
position 0 -- destroying all positional information after the prompt.
Passing ``offset=len(cache)`` selects the true absolute position.
Returns
-------
Tensor
Rotated tensor with the same shape and dtype as ``x``.
Raises
------
ValueError
If ``seq_len`` exceeds ``max_seq_len``.
"""
seq_len = x.size(2)
if offset < 0:
raise ValueError(f"offset must be >= 0, got {offset}")
if offset + seq_len > self.max_seq_len:
raise ValueError(
f"Positions [{offset}, {offset + seq_len}) exceed RoPE cache size "
f"{self.max_seq_len}. Re-instantiate with a larger max_seq_len."
)
# Retrieve cached values for this absolute position span.
# cos_cache: (seq_len, head_dim // 2)
# sin_cache: (seq_len, head_dim // 2)
cos = self.cos_cache[offset : offset + seq_len] # type: ignore[index]
sin = self.sin_cache[offset : offset + seq_len] # type: ignore[index]
# Expand to broadcast over batch and head dimensions:
# (1, 1, seq_len, head_dim // 2) → broadcasts with (B, H, T, D/2)
cos = cos.unsqueeze(0).unsqueeze(0)
sin = sin.unsqueeze(0).unsqueeze(0)
# Interleave cos and sin to match full head_dim.
# Each of (cos, sin) has shape (1, 1, T, D/2).
# We need (1, 1, T, D) by interleaving even positions with cos,
# odd positions with sin. The _rotate_half trick handles this:
#
# x_rot = x * cos_full + rotate_half(x) * sin_full
#
# where cos_full[..., 0::2] = cos and cos_full[..., 1::2] = cos,
# i.e. each cos value applies to both the even AND its paired odd slot.
# Achieved by repeating each half-dim value into both slots.
cos_full = cos.repeat_interleave(2, dim=-1) # (1, 1, T, D)
sin_full = sin.repeat_interleave(2, dim=-1) # (1, 1, T, D)
# Work in float32 for stability, then restore original dtype.
x_fp32 = x.float()
x_rot = x_fp32 * cos_full + self._rotate_half(x_fp32) * sin_full
return x_rot.to(x.dtype)
def forward(self, x: Tensor, offset: int = 0) -> Tensor:
"""Alias for :meth:`apply` to support ``nn.Sequential`` usage."""
return self.apply(x, offset=offset)
def extra_repr(self) -> str:
return (
f"head_dim={self.head_dim}, "
f"max_seq_len={self.max_seq_len}, "
f"base={self.base}"
)
# ---------------------------------------------------------------------------
# Functional helper
# ---------------------------------------------------------------------------
def apply_rope(
q: Tensor,
k: Tensor,
rope: RotaryPositionEmbedding,
offset: int = 0,
) -> Tuple[Tensor, Tensor]:
"""
Apply the same RoPE instance to both query and key tensors.
Parameters
----------
q : Tensor
Query tensor of shape ``(batch, num_heads, seq_len, head_dim)``.
k : Tensor
Key tensor of shape ``(batch, num_heads, seq_len, head_dim)``.
rope : RotaryPositionEmbedding
Pre-built RoPE module (carries the cos/sin cache on the correct device).
Returns
-------
Tuple[Tensor, Tensor]
``(q_rot, k_rot)`` — rotated queries and keys.
"""
return rope.apply(q, offset=offset), rope.apply(k, offset=offset)

View File

@@ -0,0 +1,4 @@
# configs/__init__.py
from configs.sft_config import SFTConfig, default_config, save_config, load_config
__all__ = ["SFTConfig", "default_config", "save_config", "load_config"]

View File

@@ -0,0 +1,311 @@
"""
CyberSLM SFT Configuration
===========================
Centralised dataclass for all hyperparameters and paths used
in the Supervised Instruction Fine-Tuning pipeline.
Path wiring
-----------
All default paths are resolved as absolute paths relative to the
``cyberslm_sft/`` project root so the pipeline works regardless of
where you invoke it from.
Quick setup — copy Stage 1 artifacts into the SFT project::
# 1. Copy tokenizer
cp "SLm Dataset/tokenizer/tokenizer_output/tokenizer.model" \\
"SLm Dataset/cyberslm_sft/tokenizer/tokenizer.model"
# 2. Copy pretrained weights
cp "SLm Dataset/cyberslm/checkpoints/best.pt" \\
"SLm Dataset/cyberslm_sft/checkpoints/pretrained/model.pt"
# 3. Put your instruction dataset in data/
# data/train.jsonl (required)
# data/val.jsonl (optional)
"""
from __future__ import annotations
import json
import os
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import Optional
# ---------------------------------------------------------------------------
# Project root resolution
# ---------------------------------------------------------------------------
# _PROJECT_ROOT = .../cyberslm_sft/
_PROJECT_ROOT = Path(__file__).resolve().parent.parent
# Stage 1 dataset root (one level above cyberslm_sft/) — used only as fallback
_DATASET_ROOT = _PROJECT_ROOT.parent # .../SLm Dataset/
# ---------------------------------------------------------------------------
# Model Architecture (must match pretrained checkpoint — never modify here)
# ---------------------------------------------------------------------------
@dataclass
class ModelConfig:
"""
Mirrors the pretrained CyberSLM architecture exactly.
These values are frozen; do NOT change them for SFT.
"""
vocab_size: int = 32_000
hidden_size: int = 384
num_layers: int = 12
num_heads: int = 6
head_dim: int = 64
ffn_size: int = 1024
max_seq_len: int = 2048 # MUST match the pretrained checkpoint
weight_tying: bool = True
bias: bool = False
dropout: float = 0.0
norm_eps: float = 1e-6 # RMSNorm epsilon — must match pretraining
# ---------------------------------------------------------------------------
# Tokenizer
# ---------------------------------------------------------------------------
@dataclass
class TokenizerConfig:
"""
Paths and special token ids for the SentencePiece BPE tokenizer.
model_path
Absolute path to ``tokenizer.model``.
Default: ``<project_root>/tokenizer/tokenizer.model``
(auto-falls back to Stage 1 canonical path if the local copy is absent)
"""
model_path: str = str(_PROJECT_ROOT / "tokenizer" / "tokenizer.model")
# Real SentencePiece token ids — verified against tokenizer.model
pad_id: int = 0
bos_id: int = 2 # <s> (NOT 1)
eos_id: int = 3 # </s> (NOT 2)
unk_id: int = 1 # <unk>
# ---------------------------------------------------------------------------
# Data
# ---------------------------------------------------------------------------
@dataclass
class DataConfig:
"""Dataset paths and processing knobs."""
train_path: str = str(_PROJECT_ROOT / "data" / "SFT.jsonl")
val_path: str = "" # empty = auto-split from train_path using val_split
# Maximum token length per sample (hard-truncate if exceeded).
# Must be <= ModelConfig.max_seq_len.
max_seq_len: int = 2048
# Fraction of training data to use as validation when no val_path
# is provided (ignored if val_path exists).
val_split: float = 0.05
# DataLoader workers (set 0 for debugging)
num_workers: int = 0
prefetch_factor: Optional[int] = None
# Whether to shuffle the training set each epoch
shuffle: bool = True
# Seed used for the val split and shuffling
seed: int = 42
# Cap on total samples loaded (-1 = no cap); useful for quick smoke tests
max_samples: int = -1
# ---------------------------------------------------------------------------
# Training
# ---------------------------------------------------------------------------
@dataclass
class TrainConfig:
"""
Optimizer, scheduler, and loop settings.
LR is intentionally much lower than typical pretraining (~1e-3 to 3e-4).
SFT is refinement, not relearning.
"""
# ---- Optimiser ----
learning_rate: float = 2e-5 # Peak LR after warmup
weight_decay: float = 0.01
beta1: float = 0.9
beta2: float = 0.95
eps: float = 1e-8
# ---- Gradient ----
max_grad_norm: float = 1.0
gradient_accumulation_steps: int = 4 # Effective batch = batch_size × accum
# ---- Batch ----
per_device_batch_size: int = 4
# ---- Schedule ----
num_epochs: int = 3
warmup_ratio: float = 0.03 # Fraction of total steps for warmup
lr_schedule: str = "cosine" # "cosine" | "linear" | "constant"
min_lr_ratio: float = 0.1 # min_lr = learning_rate × min_lr_ratio
# ---- Precision ----
dtype: str = "bfloat16" # bf16 on GPU; auto-disabled on CPU
# ---- Reproducibility ----
seed: int = 42
# ---- Logging ----
log_every_n_steps: int = 10
eval_every_n_steps: int = 200 # 0 = eval only at epoch end
save_every_n_steps: int = 500 # 0 = save only at epoch end
# ---- Output ----
output_dir: str = str(_PROJECT_ROOT / "checkpoints")
run_name: str = "cyberslm-instruct"
# ---- Resume ----
resume_from_checkpoint: Optional[str] = None # Path to checkpoint dir
# ---- Base model ----
pretrained_checkpoint: str = str(
_PROJECT_ROOT / "checkpoints" / "pretrained" / "model.pt"
)
# ---- Inference sanity check ----
run_inference_test: bool = True
max_new_tokens: int = 256
temperature: float = 0.7
top_p: float = 0.9
# ---------------------------------------------------------------------------
# Prompt / Template
# ---------------------------------------------------------------------------
@dataclass
class TemplateConfig:
"""
Controls which prompt format is used to wrap instruction samples.
The SFT loss is only computed on the *response* portion.
"""
# Token that separates instruction/input from the response.
# Loss masking begins immediately after this string (inclusive of the
# newline that follows it).
response_prefix: str = "### Response:\n"
# NOTE: end-of-sequence is handled as a real token id (``<eos>`` == 3),
# NOT a literal string. Appending the characters ``"</s>"`` would train
# the model to emit text that the sampler never stops on. Kept as an empty
# string so no literal EOS text is injected anywhere in the pipeline.
eos_string: str = ""
# Whether to strip leading/trailing whitespace from each field
strip_fields: bool = True
# ---------------------------------------------------------------------------
# Master Config
# ---------------------------------------------------------------------------
@dataclass
class SFTConfig:
"""
Top-level configuration. Pass this object through the entire pipeline.
Usage::
from configs.sft_config import SFTConfig, load_config, save_config
cfg = SFTConfig()
# or
cfg = load_config("my_run/sft_config.json")
"""
model: ModelConfig = field(default_factory=ModelConfig)
tokenizer: TokenizerConfig = field(default_factory=TokenizerConfig)
data: DataConfig = field(default_factory=DataConfig)
train: TrainConfig = field(default_factory=TrainConfig)
template: TemplateConfig = field(default_factory=TemplateConfig)
def __post_init__(self) -> None:
# Ensure max_seq_len for data never exceeds the model's context window
if self.data.max_seq_len > self.model.max_seq_len:
raise ValueError(
f"DataConfig.max_seq_len ({self.data.max_seq_len}) exceeds "
f"ModelConfig.max_seq_len ({self.model.max_seq_len}). "
"Truncate to the model context window or below."
)
# ---------------------------------------------------------------------------
# Serialisation helpers
# ---------------------------------------------------------------------------
def save_config(cfg: SFTConfig, path: str | Path) -> None:
"""Serialise an SFTConfig to a JSON file."""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w", encoding="utf-8") as fh:
json.dump(asdict(cfg), fh, indent=2)
def load_config(path: str | Path) -> SFTConfig:
"""
Deserialise an SFTConfig from a JSON file produced by ``save_config``.
Provides forward-compatibility: unknown keys in the file are silently
ignored so old checkpoints can still be loaded after field additions.
"""
with open(path, "r", encoding="utf-8") as fh:
raw: dict = json.load(fh)
def _filter(dc_type, d: dict) -> dict:
"""Return only keys that exist in the target dataclass."""
valid = {f.name for f in dc_type.__dataclass_fields__.values()} # type: ignore[attr-defined]
return {k: v for k, v in d.items() if k in valid}
model_cfg = ModelConfig(**_filter(ModelConfig, raw.get("model", {})))
tok_cfg = TokenizerConfig(**_filter(TokenizerConfig, raw.get("tokenizer", {})))
data_cfg = DataConfig(**_filter(DataConfig, raw.get("data", {})))
train_cfg = TrainConfig(**_filter(TrainConfig, raw.get("train", {})))
tmpl_cfg = TemplateConfig(**_filter(TemplateConfig, raw.get("template", {})))
return SFTConfig(
model=model_cfg,
tokenizer=tok_cfg,
data=data_cfg,
train=train_cfg,
template=tmpl_cfg,
)
def default_config() -> SFTConfig:
"""
Return a default SFTConfig with all paths resolved to absolute paths
inside the ``cyberslm_sft/`` project root.
Tokenizer fallback
------------------
If ``tokenizer/tokenizer.model`` doesn't exist locally, the function
automatically falls back to the canonical Stage 1 path at
``../tokenizer/tokenizer_output/tokenizer.model``.
"""
cfg = SFTConfig()
# Auto-fallback to Stage 1 tokenizer if local copy is absent
local_tok = Path(cfg.tokenizer.model_path)
canonical_tok = _DATASET_ROOT / "tokenizer" / "tokenizer_output" / "tokenizer.model"
if not local_tok.exists() and canonical_tok.exists():
import warnings
warnings.warn(
f"Local tokenizer not found at {local_tok}.\n"
f"Falling back to Stage 1 tokenizer at {canonical_tok}.\n"
"Copy it with:\n"
f" cp '{canonical_tok}' '{local_tok}'",
stacklevel=2,
)
cfg.tokenizer.model_path = str(canonical_tok)
return cfg

View File

@@ -0,0 +1,9 @@
# data/__init__.py (inference-only subset)
#
# The full package also re-exports the dataset loader, validator, collator and
# loss masking. Those are training-only and are deliberately not shipped here,
# so importing them would fail. Only the prompt formatter is needed to build a
# prompt the model recognises.
from data.prompt_formatter import PromptFormatter, Tokenizer, IGNORE_INDEX
__all__ = ["PromptFormatter", "Tokenizer", "IGNORE_INDEX"]

View File

@@ -0,0 +1,348 @@
"""
CyberSLM SFT — Prompt Formatter
================================
Converts normalised raw samples into tokenised ``(input_ids, labels)``
pairs ready for the data collator.
Tokenisation strategy (segment-based)
-------------------------------------
A sample is decomposed into an ordered list of ``(text, is_loss_target)``
segments (see ``ConversationTemplate`` for the canonical layout). Each
segment is encoded independently and the resulting token id lists are
concatenated to form ``input_ids``. ``labels`` mirror ``input_ids`` but with
non-target (prompt) positions set to ``IGNORE_INDEX`` (-100).
Why segment-based (and not char-offset) masking
------------------------------------------------
Locating the response boundary by re-encoding a prefix string and comparing
token counts is fragile: with ``add_dummy_prefix`` (and BPE merges across the
boundary) ``len(encode(prefix))`` need not equal the number of full-sequence
tokens covering that prefix, so the mask can drift by 1-2 tokens per turn.
Encoding each segment and concatenating gives an **exact** boundary because
``input_ids`` and ``labels`` are built from the very same token lists.
Special tokens (critical)
--------------------------
Special tokens are referenced by **id**, never as literal strings. The real
SentencePiece control ids are ``<bos>=2`` and ``<eos>=3``. Every assistant
response is terminated with ``eos_id`` as a genuine token and that ``eos_id``
is **left unmasked** so the model learns to stop. ``bos_id`` is prepended once
at the start of the sequence (masked). This replaces the previous, broken
approach of appending the literal characters ``"</s>"`` (which SentencePiece
encodes as ``<``,``/``,``s``,``>`` — never id 3), which prevented the model
from ever learning an end-of-sequence signal.
"""
from __future__ import annotations
import logging
from typing import List, Optional, Tuple
from configs.sft_config import SFTConfig, TemplateConfig
logger = logging.getLogger(__name__)
# Tokens with this label are excluded from the cross-entropy loss.
IGNORE_INDEX: int = -100
# One segment of the rendered prompt: (text, is_loss_target).
Segment = Tuple[str, bool]
# ---------------------------------------------------------------------------
# Segment builders (string level, EOS handled at token level downstream)
# ---------------------------------------------------------------------------
def build_alpaca_segments(
instruction: str,
output: str,
input_text: str = "",
tmpl: Optional[TemplateConfig] = None,
include_response: bool = True,
) -> List[Segment]:
"""
Build the ordered ``(text, is_loss_target)`` segments for an alpaca sample.
The response text is a loss target; everything else (instruction, input,
the response header) is masked. No EOS string is injected here — the
end-of-sequence token is appended as a real token id by
:func:`encode_segments`.
"""
if tmpl is None:
tmpl = TemplateConfig()
if tmpl.strip_fields:
instruction = instruction.strip()
input_text = input_text.strip()
output = output.strip()
if input_text:
prompt = (
f"### Instruction:\n{instruction}\n\n"
f"### Input:\n{input_text}\n\n"
f"{tmpl.response_prefix}"
)
else:
prompt = (
f"### Instruction:\n{instruction}\n\n"
f"{tmpl.response_prefix}"
)
segments: List[Segment] = [(prompt, False)]
if include_response:
segments.append((output, True))
return segments
def build_conversation_segments(
messages: List[dict],
tmpl: Optional[TemplateConfig] = None,
include_final_response: bool = True,
) -> List[Segment]:
"""
Build ordered ``(text, is_loss_target)`` segments for a multi-turn
conversation. Only assistant response bodies are loss targets.
"""
if tmpl is None:
tmpl = TemplateConfig()
def _clean(s: str) -> str:
return s.strip() if tmpl.strip_fields else s
segments: List[Segment] = []
# Optional leading system message.
body = messages
if messages and messages[0].get("role") == "system":
sys_content = _clean(messages[0]["content"])
if sys_content:
segments.append((f"System: {sys_content}\n\n", False))
body = messages[1:]
i = 0
n = len(body)
while i < n:
role = body[i]["role"]
content = _clean(body[i]["content"])
if role == "user":
segments.append((f"### User:\n{content}\n\n", False))
# Pair with the following assistant turn, if present.
if i + 1 < n and body[i + 1]["role"] == "assistant":
asst = _clean(body[i + 1]["content"])
is_last = (i + 2 >= n)
segments.append(("### Assistant:\n", False))
if not (is_last and not include_final_response):
segments.append((asst, True))
segments.append(("\n\n", False))
i += 2
else:
# Dangling user turn with no assistant reply -- the normal
# generation case. The assistant header MUST still be emitted:
# it is the cue the model was trained to continue from. Omitting
# it (the previous behaviour) handed the model a prompt shaped
# unlike anything in its training distribution.
segments.append(("### Assistant:\n", False))
i += 1
elif role == "assistant":
asst = content
segments.append(("### Assistant:\n", False))
segments.append((asst, True))
segments.append(("\n\n", False))
i += 1
else:
# Mid-conversation system message: treat as masked context.
segments.append((f"### System:\n{content}\n\n", False))
i += 1
return segments
# ---------------------------------------------------------------------------
# Tokeniser wrapper
# ---------------------------------------------------------------------------
class Tokenizer:
"""
Thin wrapper around a SentencePiece model that exposes only what the
formatter and collator need.
SentencePiece is not imported at module level so the formatter module can
be imported in unit tests without requiring sentencepiece to be installed.
"""
def __init__(self, model_path: str) -> None:
try:
import sentencepiece as spm # type: ignore
except ImportError as exc:
raise ImportError(
"sentencepiece is required for the tokenizer. "
"Install with: pip install sentencepiece"
) from exc
self._sp = spm.SentencePieceProcessor()
self._sp.Load(model_path)
self.bos_id: int = self._sp.bos_id()
self.eos_id: int = self._sp.eos_id()
self.pad_id: int = self._sp.pad_id()
self.vocab_size: int = self._sp.GetPieceSize()
# Fail loudly if the tokenizer's control ids do not match the ids the
# rest of the pipeline (and the base model config) assume. A silent
# mismatch here corrupts training and breaks generation stopping.
if self.eos_id < 0:
raise ValueError(
"Tokenizer has no EOS id (eos_id < 0). The SFT pipeline "
"requires a real end-of-sequence token."
)
# ------------------------------------------------------------------
def encode(
self,
text: str,
add_bos: bool = False,
add_eos: bool = False,
) -> List[int]:
ids: List[int] = self._sp.Encode(text, out_type=int)
if add_bos and self.bos_id >= 0:
ids = [self.bos_id] + ids
if add_eos and self.eos_id >= 0:
ids = ids + [self.eos_id]
return ids
def decode(self, ids: List[int]) -> str:
return self._sp.Decode(ids)
def __len__(self) -> int:
return self.vocab_size
# ---------------------------------------------------------------------------
# Segment → token id encoding
# ---------------------------------------------------------------------------
def encode_segments(
segments: List[Segment],
tokenizer: Tokenizer,
max_seq_len: int,
add_bos: bool = True,
) -> Tuple[List[int], List[int]]:
"""
Encode ``(text, is_loss_target)`` segments into ``(input_ids, labels)``.
* ``input_ids`` is the concatenation of each segment's token ids.
* A real ``eos_id`` token is appended immediately after every loss-target
(assistant) segment and is itself a loss target — the model must learn
to emit it.
* ``bos_id`` is prepended once at the start (masked) when ``add_bos``.
* ``labels`` equal ``input_ids`` on target positions and ``IGNORE_INDEX``
elsewhere.
The sequence is truncated to ``max_seq_len`` tokens.
"""
input_ids: List[int] = []
labels: List[int] = []
if add_bos and tokenizer.bos_id is not None and tokenizer.bos_id >= 0:
input_ids.append(tokenizer.bos_id)
labels.append(IGNORE_INDEX)
for text, is_target in segments:
ids = tokenizer.encode(text)
input_ids.extend(ids)
labels.extend(ids if is_target else [IGNORE_INDEX] * len(ids))
if is_target:
# Terminate the assistant turn with a learned EOS token.
input_ids.append(tokenizer.eos_id)
labels.append(tokenizer.eos_id)
return input_ids[:max_seq_len], labels[:max_seq_len]
# ---------------------------------------------------------------------------
# Public API — PromptFormatter
# ---------------------------------------------------------------------------
class PromptFormatter:
"""
Converts a normalised raw sample into a tokenised ``(input_ids, labels)``
pair, applying loss-masking so only assistant responses are trained on.
Parameters
----------
cfg:
Master ``SFTConfig`` (for template and data settings).
tokenizer:
Initialised ``Tokenizer`` instance.
"""
def __init__(self, cfg: SFTConfig, tokenizer: Tokenizer) -> None:
self.cfg = cfg
self.tokenizer = tokenizer
self.tmpl = cfg.template
self.max_len = cfg.data.max_seq_len
# ------------------------------------------------------------------
def format(self, sample: dict) -> Optional[Tuple[List[int], List[int]]]:
"""
Format a single normalised sample into ``(input_ids, labels)``.
Returns ``None`` when the sample produces no trainable (non-masked)
tokens — e.g. the response was truncated away.
"""
if "messages" in sample:
segments = build_conversation_segments(
sample["messages"], tmpl=self.tmpl, include_final_response=True
)
else:
segments = build_alpaca_segments(
instruction=sample["instruction"],
output=sample["output"],
input_text=sample.get("input", ""),
tmpl=self.tmpl,
include_response=True,
)
# Nothing to train on if there are no target segments at all.
if not any(is_target for _, is_target in segments):
logger.debug("No assistant/response segments — skipping sample")
return None
input_ids, labels = encode_segments(
segments, self.tokenizer, self.max_len, add_bos=True
)
if not input_ids or not any(l != IGNORE_INDEX for l in labels):
logger.debug(
"No response tokens remain after truncation at %d — skipping",
self.max_len,
)
return None
return input_ids, labels
# ------------------------------------------------------------------
def format_for_inference(self, sample: dict) -> List[int]:
"""
Format a sample for *generation*: prompt only, no response body and no
trailing EOS, with ``bos_id`` prepended.
"""
if "messages" in sample:
segments = build_conversation_segments(
sample["messages"], tmpl=self.tmpl, include_final_response=False
)
else:
segments = build_alpaca_segments(
instruction=sample["instruction"],
output=sample.get("output", ""),
input_text=sample.get("input", ""),
tmpl=self.tmpl,
include_response=False,
)
input_ids, _ = encode_segments(
segments, self.tokenizer, self.max_len, add_bos=True
)
return input_ids

View File

@@ -0,0 +1,3 @@
from .cyberslm import CyberSLM
__all__ = ["CyberSLM"]

View File

@@ -0,0 +1,40 @@
import sys
from pathlib import Path
# Add the root directory (which contains 'cyberslm') to sys.path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent.parent))
from cyberslm.model.model import CyberSLM as Stage1CyberSLM
from cyberslm.model.config import CyberSLMConfig
class CyberSLM(Stage1CyberSLM):
"""
Adapter class to bridge SFT's ModelConfig to Stage 1's CyberSLMConfig,
allowing us to use the original Stage 1 model architecture verbatim.
"""
def __init__(self, cfg):
# Convert SFT ModelConfig to Stage 1 CyberSLMConfig
stage1_cfg = CyberSLMConfig(
vocab_size=cfg.vocab_size,
hidden_dim=cfg.hidden_size,
num_layers=cfg.num_layers,
num_heads=cfg.num_heads,
head_dim=cfg.head_dim,
ffn_hidden_dim=cfg.ffn_size,
max_seq_len=cfg.max_seq_len,
tie_weights=cfg.weight_tying,
bias=cfg.bias,
dropout=cfg.dropout,
norm_eps=cfg.norm_eps,
)
super().__init__(stage1_cfg)
def forward(self, input_ids, attention_mask=None, **kwargs):
"""
Wrapper around Stage 1's forward pass.
Stage 1 returns (logits, all_attn_weights); the SFT trainer expects
just the logits Tensor. The key-padding ``attention_mask`` (1=keep,
0=pad) is forwarded so right-padded batches do not corrupt real tokens.
"""
logits, _ = super().forward(input_ids, attention_mask=attention_mask)
return logits

10
generation_config.json Normal file
View File

@@ -0,0 +1,10 @@
{
"_from_model_config": true,
"bos_token_id": 2,
"eos_token_id": 3,
"output_attentions": false,
"output_hidden_states": false,
"pad_token_id": 0,
"transformers_version": "5.15.1",
"use_cache": true
}

118
infer_chat.py Normal file
View File

@@ -0,0 +1,118 @@
"""
Instruction-tuned model inference — question answering.
python Final/infer_chat.py --prompt "What is SQL injection?"
python Final/infer_chat.py --interactive
The prompt is built with the SAME formatter used during fine-tuning, so the
model sees exactly the token sequence it was trained on. Hand-assembling the
prompt string instead produces different token ids at every segment boundary
(SentencePiece prepends a word-start marker per encode call) and the model then
sees something it was never trained on.
Expect correctly-shaped answers with unreliable facts: this is a 33.5M-parameter
model. See runs/reports/FINAL_REPORT.md for measured behaviour.
"""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
import torch
_HERE = Path(__file__).resolve().parent
for _p in (_HERE, _HERE / "cyberslm_sft"):
if str(_p) not in sys.path:
sys.path.insert(0, str(_p))
from configs.sft_config import default_config # noqa: E402
from data.prompt_formatter import PromptFormatter, Tokenizer # noqa: E402
from model.cyberslm import CyberSLM as SFTModel # noqa: E402
def main() -> int:
ap = argparse.ArgumentParser(description="CyberSLM instruct (question answering)")
ap.add_argument("--prompt", "-p", default="What is SQL injection and how do I prevent it?")
ap.add_argument("--interactive", "-i", action="store_true")
ap.add_argument("--checkpoint", "-c", default=str(_HERE / "models" / "instruct.pt"))
ap.add_argument("--tokenizer", default=str(_HERE / "tokenizer" / "tokenizer.model"))
ap.add_argument("--max-new-tokens", "-m", type=int, default=200)
ap.add_argument("--temperature", "-t", type=float, default=0.0,
help="0 = greedy/deterministic (recommended for this model)")
ap.add_argument("--top-k", type=int, default=50)
ap.add_argument("--top-p", type=float, default=0.9)
ap.add_argument("--repetition-penalty", type=float, default=1.1)
ap.add_argument("--device", default=None)
args = ap.parse_args()
device = torch.device(args.device) if args.device else torch.device(
"cuda" if torch.cuda.is_available() else "cpu")
ckpt = Path(args.checkpoint)
if not ckpt.exists():
print(f"Checkpoint not found: {ckpt}", file=sys.stderr)
return 1
cfg = default_config()
cfg.tokenizer.model_path = args.tokenizer
cfg.model.max_seq_len = 2048
cfg.data.max_seq_len = 2048
tok = Tokenizer(cfg.tokenizer.model_path)
fmt = PromptFormatter(cfg=cfg, tokenizer=tok)
model = SFTModel(cfg.model)
state = torch.load(ckpt, map_location=device, weights_only=False)
if isinstance(state, dict) and "model_state" in state:
state = state["model_state"]
model.load_state_dict(state)
model.to(device).eval()
n = sum(p.numel() for p in model.parameters())
print(f"model : {ckpt.name} ({n:,} params)")
print(f"context: {cfg.model.max_seq_len} device: {device} "
f"decoding: {'greedy' if args.temperature == 0 else f'T={args.temperature}'}")
def answer(question: str) -> None:
ids = fmt.format_for_inference({"messages": [{"role": "user", "content": question}]})
x = torch.tensor([ids], dtype=torch.long, device=device)
t0 = time.perf_counter()
out = model.generate(
x, max_new_tokens=args.max_new_tokens, temperature=args.temperature,
top_k=args.top_k, top_p=args.top_p,
repetition_penalty=args.repetition_penalty, eos_id=tok.eos_id,
)
dt = time.perf_counter() - t0
new = out[0, len(ids):].tolist()
stopped = tok.eos_id in new
if stopped:
new = new[: new.index(tok.eos_id)]
print("\n" + "-" * 66)
print(tok.decode(new).strip() or "(empty)")
print("-" * 66)
print(f"{len(new)} tokens in {dt:.2f}s ({len(new)/dt if dt else 0:.1f} tok/s), "
f"{'stopped on EOS' if stopped else 'hit token limit'}\n")
if args.interactive:
print("\nInstruct model - ask a question. ('exit' to quit)")
while True:
try:
q = input("\nYou: ").strip()
except (EOFError, KeyboardInterrupt):
print("\nBye."); break
if not q:
continue
if q.lower() in {"exit", "quit", "q"}:
print("Bye."); break
answer(q)
return 0
answer(args.prompt)
return 0
if __name__ == "__main__":
sys.exit(main())

3
model.safetensors Normal file
View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:a29dca58c1ac9b190015c5fce60425830a2ee0009117969cefcb9d7fab569a06
size 134137152

3
models/instruct.pt Normal file
View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:baf9e480ae7ce58282c74703b5eeac3d80403ec4cbe102c179382cda2d4926b6
size 134160315

6
special_tokens_map.json Normal file
View File

@@ -0,0 +1,6 @@
{
"bos_token": "<bos>",
"eos_token": "<eos>",
"unk_token": "<unk>",
"pad_token": "<pad>"
}

279575
tokenizer.json Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:0c1a2ad32fd7d576e42bad69c99312896a077913a543a25af73831e3266a68e6
size 735456

9
tokenizer_config.json Normal file
View File

@@ -0,0 +1,9 @@
{
"tokenizer_class": "PreTrainedTokenizerFast",
"model_max_length": 2048,
"bos_token": "<bos>",
"eos_token": "<eos>",
"unk_token": "<unk>",
"pad_token": "<pad>",
"clean_up_tokenization_spaces": false
}