echo / MVP /distill.py
void0x14
feat: multimodal model anahtar teslim + test + rapor
598b018 unverified
Raw History Blame Contribute Delete
12.1 kB
from __future__ import annotations
import argparse
import json
import math
import time
from dataclasses import dataclass
from pathlib import Path
import torch
import torch.nn.functional as F
from huggingface_hub import hf_hub_download
from safetensors.torch import save_file
from transformers import PreTrainedModel
from torch.optim import AdamW
from torch.utils.data import DataLoader, Dataset
from transformers import AutoConfig, AutoModel, AutoModelForCausalLM, AutoTokenizer
@dataclass
class DistillConfig:
teacher_path: str
student_path: str
output_dir: str
dataset_path: str = ""
dataset_split: str = "train"
seq_len: int = 128
batch_size: int = 1
grad_accum: int = 4
lr: float = 3e-4
warmup_ratio: float = 0.05
max_steps: int = 1000
log_every: int = 25
save_every: int = 250
resume_from: str = ""
ce_loss_weight: float = 1.0
mse_loss_weight: float = 1.0
max_grad_norm: float = 1.0
seed: int = 42
class TokenizedDataset(Dataset):
def __init__(self, token_ids: list[int], seq_len: int):
self.seq_len = seq_len
self.examples = []
for i in range(0, len(token_ids) - seq_len - 1, seq_len):
chunk = token_ids[i : i + seq_len + 1]
if len(chunk) == seq_len + 1:
self.examples.append(torch.tensor(chunk, dtype=torch.long))
def __len__(self) -> int:
return len(self.examples)
def __getitem__(self, idx: int) -> torch.Tensor:
return self.examples[idx]
def load_teacher(path: str, dtype: torch.dtype) -> tuple:
model = AutoModel.from_pretrained(path, trust_remote_code=True, dtype=dtype)
model.eval()
for p in model.parameters():
p.requires_grad = False
embed_weight = model.language_model.embed_tokens.weight
return model, embed_weight
def load_student(path: str, dtype: torch.dtype) -> tuple:
config = AutoConfig.from_pretrained(path, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(
path, trust_remote_code=True, dtype=dtype
)
model.train()
return model, config
def tokenize_dataset(dataset_path: str, tokenizer, max_tokens: int = 500000) -> list[int]:
"""Read a raw text file and tokenize it."""
path = Path(dataset_path)
if not path.exists():
raise FileNotFoundError(f"Dataset not found: {dataset_path}")
text = path.read_text(encoding="utf-8")
all_ids: list[int] = []
for paragraph in text.split("\n\n"):
paragraph = paragraph.strip()
if not paragraph:
continue
ids = tokenizer.encode(paragraph, add_special_tokens=False)
all_ids.extend(ids)
all_ids.append(tokenizer.eos_token_id)
if len(all_ids) >= max_tokens:
break
return all_ids[:max_tokens]
def compute_teacher_outputs(teacher, embed_weight, input_ids: torch.Tensor):
with torch.no_grad():
out = teacher(
input_ids=input_ids,
output_hidden_states=True,
use_cache=False,
)
hidden = out.hidden_states[-1]
logits = hidden @ embed_weight.T
return logits, hidden
def compute_student_outputs(student, input_ids: torch.Tensor):
out = student(
input_ids=input_ids,
output_hidden_states=True,
use_cache=False,
)
return out.logits, out.hidden_states[-1]
def distillation_loss(
student_logits: torch.Tensor,
teacher_logits: torch.Tensor,
student_hidden: torch.Tensor,
teacher_hidden: torch.Tensor,
ce_weight: float,
mse_weight: float,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
shift_student_logits = student_logits[:, :-1, :].contiguous()
shift_teacher_logits = teacher_logits[:, :-1, :].contiguous()
teacher_probs = F.softmax(shift_teacher_logits, dim=-1)
student_log_probs = F.log_softmax(shift_student_logits, dim=-1)
ce_loss = -(teacher_probs * student_log_probs).sum(dim=-1).mean()
shift_student_hidden = student_hidden[:, :-1, :].contiguous()
shift_teacher_hidden = teacher_hidden[:, :-1, :].contiguous()
mse_loss = F.mse_loss(shift_student_hidden, shift_teacher_hidden)
total = ce_weight * ce_loss + mse_weight * mse_loss
return total, ce_loss.detach(), mse_loss.detach()
def get_cosine_schedule_with_warmup(optimizer, warmup_steps: int, total_steps: int):
def lr_lambda(current_step: int) -> float:
if current_step < warmup_steps:
return float(current_step) / float(max(1, warmup_steps))
progress = float(current_step - warmup_steps) / float(max(1, total_steps - warmup_steps))
return max(0.0, 0.5 * (1.0 + math.cos(math.pi * progress)))
return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
def save_checkpoint(student, student_config, output_dir: str, step: int):
out = Path(output_dir)
out.mkdir(parents=True, exist_ok=True)
state_dict = {}
for k, v in student.state_dict().items():
if k == "lm_head.weight" and "model.embed_tokens.weight" in state_dict:
state_dict[k] = state_dict["model.embed_tokens.weight"].clone()
else:
state_dict[k] = v.cpu()
save_file(state_dict, str(out / "model.safetensors"), metadata={"format": "pt"})
config_path = out / "config.json"
config_path.write_text(json.dumps(student_config.to_dict(), indent=2, sort_keys=True) + "\n", encoding="utf-8")
(out / "train_state.json").write_text(json.dumps({"step": step}) + "\n", encoding="utf-8")
print(f" Checkpoint saved at step {step}: {out}", flush=True)
def run_distill(args: DistillConfig) -> None:
torch.manual_seed(args.seed)
device = torch.device("cpu")
dtype = torch.float32
print(f"Loading teacher from {args.teacher_path}...", flush=True)
teacher, embed_weight = load_teacher(args.teacher_path, dtype)
teacher_params = sum(p.numel() for p in teacher.parameters())
print(f" Teacher loaded: {teacher_params / 1e6:.1f}M params", flush=True)
print(f"Loading student from {args.student_path}...", flush=True)
student, student_config = load_student(args.student_path, dtype)
student_params = sum(p.numel() for p in student.parameters())
print(f" Student loaded: {student_params / 1e6:.1f}M params", flush=True)
print("Loading tokenizer...", flush=True)
tokenizer = AutoTokenizer.from_pretrained(args.student_path, trust_remote_code=True)
print(f"Tokenizing dataset ({args.dataset_path})...", flush=True)
all_ids = tokenize_dataset(args.dataset_path, tokenizer)
print(f" Total tokens: {len(all_ids)}", flush=True)
dataset = TokenizedDataset(all_ids, args.seq_len)
dataloader = DataLoader(dataset, batch_size=args.batch_size, shuffle=True, drop_last=True)
print(f" Dataset size: {len(dataset)} examples", flush=True)
optimizer = AdamW(
[p for p in student.parameters() if p.requires_grad],
lr=args.lr,
weight_decay=0.01,
)
scheduler = get_cosine_schedule_with_warmup(
optimizer,
warmup_steps=int(args.max_steps * args.warmup_ratio),
total_steps=args.max_steps,
)
print(f"Starting distillation: {args.max_steps} steps, lr={args.lr}", flush=True)
print(f" Loss weights: CE={args.ce_loss_weight}, MSE={args.mse_loss_weight}", flush=True)
print(f" Batch size={args.batch_size}, grad_accum={args.grad_accum}, seq_len={args.seq_len}", flush=True)
start_step = 0
if args.resume_from:
ckpt_path = Path(args.resume_from)
ckpt_file = ckpt_path / "model.safetensors"
state_file = ckpt_path / "train_state.json"
if ckpt_file.exists():
from safetensors.torch import load_file
ckpt_state = load_file(str(ckpt_file))
missing, unexpected = student.load_state_dict(ckpt_state, strict=False)
if state_file.exists():
start_step = json.loads(state_file.read_text())["step"]
print(f" Resumed from {ckpt_path} at step {start_step}: missing={len(missing)}, unexpected={len(unexpected)}", flush=True)
else:
print(f" Warning: checkpoint not found at {ckpt_file}, starting from scratch", flush=True)
step = start_step
running_ce = 0.0
running_mse = 0.0
running_total = 0.0
start_time = time.time()
steps_done = 0
student.train()
while step < args.max_steps:
for batch in dataloader:
if step >= args.max_steps:
break
input_ids = batch.to(device)
teacher_logits, teacher_hidden = compute_teacher_outputs(teacher, embed_weight, input_ids)
student_logits, student_hidden = compute_student_outputs(student, input_ids)
loss, ce_loss, mse_loss = distillation_loss(
student_logits, teacher_logits, student_hidden, teacher_hidden,
args.ce_loss_weight, args.mse_loss_weight,
)
loss = loss / args.grad_accum
loss.backward()
running_ce += ce_loss.item()
running_mse += mse_loss.item()
running_total += loss.item() * args.grad_accum
if (step + 1) % args.grad_accum == 0:
torch.nn.utils.clip_grad_norm_(student.parameters(), args.max_grad_norm)
optimizer.step()
scheduler.step()
optimizer.zero_grad()
if (step + 1) % args.log_every == 0:
elapsed = time.time() - start_time
avg_ce = running_ce / args.log_every
avg_mse = running_mse / args.log_every
avg_total = running_total / args.log_every
lr_now = scheduler.get_last_lr()[0]
steps_per_sec = (step + 1) / elapsed
eta = (args.max_steps - step - 1) / steps_per_sec if steps_per_sec > 0 else 0
print(
f" Step {step + 1}/{args.max_steps} | "
f"total={avg_total:.4f} ce={avg_ce:.4f} mse={avg_mse:.6f} | "
f"lr={lr_now:.2e} | {elapsed:.0f}s elapsed, ~{eta:.0f}s remaining",
flush=True,
)
running_ce = 0.0
running_mse = 0.0
running_total = 0.0
if (step + 1) % args.save_every == 0:
save_checkpoint(student, student_config, args.output_dir, step + 1)
step += 1
final_dir = Path(args.output_dir) / "final"
save_checkpoint(student, student_config, str(final_dir), step)
print(f"\nDistillation complete. Final model saved to {final_dir}")
print(f"Total time: {time.time() - start_time:.0f}s")
def _main() -> None:
parser = argparse.ArgumentParser(description="Knowledge distillation: Qwen3.5 teacher -> pruned student")
parser.add_argument("--teacher-path", required=True)
parser.add_argument("--student-path", required=True)
parser.add_argument("--output-dir", required=True)
parser.add_argument("--dataset-path", required=True, help="Path to raw text file")
parser.add_argument("--seq-len", type=int, default=256)
parser.add_argument("--batch-size", type=int, default=1)
parser.add_argument("--grad-accum", type=int, default=4)
parser.add_argument("--lr", type=float, default=3e-4)
parser.add_argument("--warmup-ratio", type=float, default=0.05)
parser.add_argument("--max-steps", type=int, default=1000)
parser.add_argument("--log-every", type=int, default=25)
parser.add_argument("--save-every", type=int, default=250)
parser.add_argument("--ce-loss-weight", type=float, default=1.0)
parser.add_argument("--mse-loss-weight", type=float, default=1.0)
parser.add_argument("--max-grad-norm", type=float, default=1.0)
parser.add_argument("--resume-from", type=str, default="", help="Path to checkpoint dir to resume from")
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
run_distill(DistillConfig(**{k.replace("-", "_"): v for k, v in vars(args).items()}))
if __name__ == "__main__":
_main()