初始化项目,由ModelHub XC社区提供模型
Model: sabari2005/cyberslm-instruct Source: Original Platform
This commit is contained in:
4
cyberslm_sft/configs/__init__.py
Normal file
4
cyberslm_sft/configs/__init__.py
Normal 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"]
|
||||
311
cyberslm_sft/configs/sft_config.py
Normal file
311
cyberslm_sft/configs/sft_config.py
Normal 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
|
||||
9
cyberslm_sft/data/__init__.py
Normal file
9
cyberslm_sft/data/__init__.py
Normal 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"]
|
||||
348
cyberslm_sft/data/prompt_formatter.py
Normal file
348
cyberslm_sft/data/prompt_formatter.py
Normal 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
|
||||
3
cyberslm_sft/model/__init__.py
Normal file
3
cyberslm_sft/model/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
from .cyberslm import CyberSLM
|
||||
|
||||
__all__ = ["CyberSLM"]
|
||||
40
cyberslm_sft/model/cyberslm.py
Normal file
40
cyberslm_sft/model/cyberslm.py
Normal 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
|
||||
Reference in New Issue
Block a user