Files
Bharat-Tiny-LLM-v2/brahmi_continued_pretrain.ipynb
ModelHub XC b416fc0c1e 初始化项目,由ModelHub XC社区提供模型
Model: eulogik/Bharat-Tiny-LLM-v2
Source: Original Platform
2026-07-25 18:04:11 +08:00

395 lines
14 KiB
Plaintext
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

{
"cells": [
{
"cell_type": "markdown",
"metadata": {"id": "heading"},
"source": [
"# Brahmi Embedding Warmup — Aggressive Round 2\n",
"\n",
"**Goal**: Train the 300 new Devanagari embeddings using ALL 435K training rows.\n",
"Round 1 (150 steps, 2000 chunks) produced garbled Hindi. This round uses:\n",
"- Full corpus (435K rows → ~200K+ chunks)\n",
"- Higher LR (1e-3)\n",
"- Cosine schedule\n",
"- Only chunks containing new tokens (efficient training)\n",
"\n",
"**Hardware**: T4 GPU (Colab free tier). ~2-3 hours.\n",
"\n",
"**Before starting**: Paste your HF_TOKEN below."
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "install"},
"outputs": [],
"source": [
"!pip install -q transformers torch datasets accelerate huggingface_hub\n",
"\n",
"import json, os, math, time, random\n",
"import numpy as np\n",
"import torch\n",
"import torch.nn as nn\n",
"from torch.utils.data import DataLoader\n",
"from transformers import AutoTokenizer, AutoModelForCausalLM\n",
"\n",
"HF_TOKEN = \"hf_YOUR_TOKEN_HERE\"\n",
"\n",
"# Mount Google Drive for persistence across sessions\n",
"from google.colab import drive\n",
"drive.mount('/content/drive')\n",
"DRIVE_DIR = '/content/drive/MyDrive/brahmi_training'\n",
"os.makedirs(DRIVE_DIR, exist_ok=True)\n",
"print(f'Saving to {DRIVE_DIR}')\n",
"\n",
"from huggingface_hub import login\n",
"login(token=HF_TOKEN)\n",
"\n",
"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
"print(f'Device: {device}')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "load"},
"outputs": [],
"source": [
"MODEL_ID = 'eulogik/Bharat-Tiny-LLM-v2'\n",
"\n",
"tok = AutoTokenizer.from_pretrained(MODEL_ID, token=HF_TOKEN)\n",
"if tok.pad_token is None:\n",
" tok.pad_token = tok.eos_token\n",
"\n",
"model = AutoModelForCausalLM.from_pretrained(\n",
" MODEL_ID,\n",
" torch_dtype=torch.bfloat16,\n",
" device_map='auto',\n",
" token=HF_TOKEN,\n",
")\n",
"\n",
"new_ids = sorted(\n",
" tid for tid, t in tok.added_tokens_decoder.items()\n",
" if not str(t).startswith('<') and not getattr(t, 'special', False)\n",
")\n",
"print(f'Found {len(new_ids)} new tokens (IDs {new_ids[0]}{new_ids[-1]})')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "freeze"},
"outputs": [],
"source": [
"for p in model.parameters():\n",
" p.requires_grad = False\n",
"model.model.embed_tokens.weight.requires_grad_(True)\n",
"model.lm_head.weight.requires_grad_(True)\n",
"\n",
"class LearnableRows(nn.Module):\n",
" def __init__(self, embed, lmhead):\n",
" super().__init__()\n",
" self.embed = nn.Parameter(embed)\n",
" self.lmhead = nn.Parameter(lmhead)\n",
"\n",
"hidden = model.config.hidden_size\n",
"embed_w = model.model.embed_tokens.weight\n",
"lmhead_w = model.lm_head.weight\n",
"\n",
"learner = LearnableRows(\n",
" embed_w.data[new_ids].clone(),\n",
" lmhead_w.data[new_ids].clone()\n",
").to(device)\n",
"\n",
"n_params = len(new_ids) * hidden * 2\n",
"print(f'Training {n_params:,} params (0.04% of {sum(p.numel() for p in model.parameters()):,})')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "data_load"},
"outputs": [],
"source": [
"import requests, gzip, shutil\n",
"\n",
"DATA_URL = 'https://huggingface.co/eulogik/Bharat-Tiny-LLM-v2/resolve/main/train_gold_v3.jsonl.gz'\n",
"DATA_FILE = '/content/train_gold_v3.jsonl'\n",
"\n",
"if not os.path.exists(DATA_FILE):\n",
" print('Downloading...')\n",
" r = requests.get(DATA_URL, stream=True,\n",
" headers={'Authorization': f'Bearer {HF_TOKEN}'})\n",
" r.raise_for_status()\n",
" with open('/content/data.gz', 'wb') as f:\n",
" shutil.copyfileobj(r.raw, f)\n",
" with gzip.open('/content/data.gz', 'rb') as gz, open(DATA_FILE, 'wb') as f:\n",
" shutil.copyfileobj(gz, f)\n",
" os.remove('/content/data.gz')\n",
" print(f'Ready! {os.path.getsize(DATA_FILE) / 1e6:.0f} MB')\n",
"\n",
"from datasets import load_dataset\n",
"dataset = load_dataset('text', data_files=DATA_FILE, split='train')\n",
"print(f'Loaded {len(dataset)} rows')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "tokenize"},
"outputs": [],
"source": [
"CHUNKS_FILE = os.path.join(DRIVE_DIR, 'all_chunks.json')\n",
"\n",
"if os.path.exists(CHUNKS_FILE):\n",
" print('Loading cached chunks from Drive...')\n",
" with open(CHUNKS_FILE) as f:\n",
" all_chunks = json.load(f)\n",
" random.shuffle(all_chunks)\n",
" print(f'Loaded {len(all_chunks)} chunks')\n",
"else:\n",
" new_id_set = set(new_ids)\n",
" all_chunks = []\n",
" skipped = 0\n",
"\n",
" for i, example in enumerate(dataset):\n",
" text = example.get('text', '')\n",
" if isinstance(text, str) and text.startswith('{'):\n",
" try:\n",
" data = json.loads(text)\n",
" text = ' '.join(m['content'] for m in data.get('messages', []))\n",
" except:\n",
" pass\n",
" ids = tok.encode(text)\n",
" for j in range(0, len(ids), 256):\n",
" chunk = ids[j:j+256]\n",
" if len(chunk) >= 10:\n",
" if any(tid in new_id_set for tid in chunk):\n",
" all_chunks.append(chunk)\n",
" else:\n",
" skipped += 1\n",
" if i % 50000 == 0 and i > 0:\n",
" print(f' Processed {i}/{len(dataset)} rows...')\n",
"\n",
" random.shuffle(all_chunks)\n",
" with open(CHUNKS_FILE, 'w') as f:\n",
" json.dump(all_chunks, f)\n",
" print(f'\\nKept {len(all_chunks)} chunks (skipped {skipped})')\n",
" print(f'Saved to Drive')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "dataloader"},
"outputs": [],
"source": [
"def collate(batch):\n",
" mx = max(len(x) for x in batch)\n",
" pad = torch.zeros(len(batch), mx, dtype=torch.long)\n",
" for i, x in enumerate(batch):\n",
" pad[i, :len(x)] = torch.tensor(x, dtype=torch.long)\n",
" return pad.to(device)\n",
"\n",
"split = int(len(all_chunks) * 0.95)\n",
"train_chunks = all_chunks[:split]\n",
"val_chunks = all_chunks[split:]\n",
"\n",
"BATCH_SIZE = 4\n",
"train_loader = DataLoader(train_chunks, batch_size=BATCH_SIZE, shuffle=True, collate_fn=collate)\n",
"val_loader = DataLoader(val_chunks, batch_size=BATCH_SIZE, collate_fn=collate)\n",
"\n",
"print(f'Train: {len(train_chunks)} chunks ({len(train_loader)} batches)')\n",
"print(f'Val: {len(val_chunks)} chunks ({len(val_loader)} batches)')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "init_val"},
"outputs": [],
"source": [
"NUM_STEPS = 3000\n",
"LR = 1e-3\n",
"CKPT_DIR = os.path.join(DRIVE_DIR, 'checkpoints')\n",
"os.makedirs(CKPT_DIR, exist_ok=True)\n",
"\n",
"optim = torch.optim.AdamW(learner.parameters(), lr=LR, weight_decay=0.01)\n",
"scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n",
" optim, T_max=NUM_STEPS, eta_min=1e-5\n",
")\n",
"\n",
"# Resume from checkpoint if exists\n",
"step = 0\n",
"best_loss = float('inf')\n",
"init_loss = None\n",
"resume_path = os.path.join(CKPT_DIR, 'state.json')\n",
"learner_path = os.path.join(CKPT_DIR, 'learner.pt')\n",
"\n",
"if os.path.exists(resume_path) and os.path.exists(learner_path):\n",
" with open(resume_path) as f:\n",
" state = json.load(f)\n",
" step = state['step']\n",
" best_loss = state['best_loss']\n",
" init_loss = state['init_loss']\n",
" learner.load_state_dict(torch.load(learner_path, map_location=device))\n",
" # Fast-forward scheduler\n",
" for _ in range(step):\n",
" scheduler.step()\n",
" print(f'Resumed from step {step}, best loss {best_loss:.4f}')\n",
"else:\n",
" # Initial validation loss\n",
" model.eval()\n",
" init_losses = []\n",
" with torch.no_grad():\n",
" for i, batch in enumerate(val_loader):\n",
" init_losses.append(model(batch, labels=batch).loss.item())\n",
" if i >= 20: break\n",
" init_loss = np.mean(init_losses)\n",
" best_loss = init_loss\n",
" print(f'Initial val loss: {init_loss:.4f} (PPL: {math.exp(init_loss):.1f})')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "train_loop"},
"outputs": [],
"source": [
"model.train()\n",
"start_t = time.time()\n",
"log_every = 100\n",
"\n",
"for epoch in range(10):\n",
" for batch in train_loader:\n",
" if step >= NUM_STEPS:\n",
" break\n",
"\n",
" # Copy learner → model\n",
" with torch.no_grad():\n",
" embed_w.data[new_ids] = learner.embed.data.to(embed_w.device)\n",
" lmhead_w.data[new_ids] = learner.lmhead.data.to(lmhead_w.device)\n",
"\n",
" out = model(batch, labels=batch)\n",
" loss = out.loss\n",
" loss.backward()\n",
"\n",
" # Copy grads → learner\n",
" with torch.no_grad():\n",
" learner.embed.grad = embed_w.grad[new_ids].clone().to(learner.embed.device)\n",
" learner.lmhead.grad = lmhead_w.grad[new_ids].clone().to(learner.lmhead.device)\n",
"\n",
" torch.nn.utils.clip_grad_norm_(learner.parameters(), 1.0)\n",
" optim.step()\n",
" scheduler.step()\n",
" optim.zero_grad()\n",
" model.zero_grad()\n",
" step += 1\n",
"\n",
" if step % log_every == 0:\n",
" model.eval()\n",
" vlosses = []\n",
" with torch.no_grad():\n",
" for vb in val_loader:\n",
" vlosses.append(model(vb, labels=vb).loss.item())\n",
" if len(vlosses) >= 20: break\n",
" avg_vloss = np.mean(vlosses)\n",
" impr = (init_loss - avg_vloss) / init_loss * 100\n",
" lr_now = scheduler.get_last_lr()[0]\n",
" elapsed = time.time() - start_t\n",
" print(f'Step {step}/{NUM_STEPS} | '\n",
" f'Tr:{loss.item():.4f} | '\n",
" f'Val:{avg_vloss:.4f} (PPL:{math.exp(avg_vloss):.1f}) | '\n",
" f'{impr:+.1f}% | '\n",
" f'LR:{lr_now:.2e} | '\n",
" f'{step/elapsed:.2f} it/s')\n",
" if avg_vloss < best_loss:\n",
" best_loss = avg_vloss\n",
" with torch.no_grad():\n",
" embed_w.data[new_ids] = learner.embed.data.to(embed_w.device)\n",
" lmhead_w.data[new_ids] = learner.lmhead.data.to(lmhead_w.device)\n",
" model.save_pretrained('brahmi-best')\n",
" tok.save_pretrained('brahmi-best')\n",
" # Save checkpoint for resume\n",
" torch.save(learner.state_dict(), learner_path)\n",
" with open(resume_path, 'w') as f:\n",
" json.dump({'step': step, 'best_loss': best_loss, 'init_loss': init_loss}, f)\n",
" model.train()\n",
"\n",
" if step >= NUM_STEPS:\n",
" break\n",
"\n",
"elapsed = time.time() - start_t\n",
"final_impr = (init_loss - best_loss) / init_loss * 100\n",
"print(f'\\nDone! {step} steps in {elapsed:.0f}s ({step/elapsed:.2f} it/s)')\n",
"print(f'Init loss: {init_loss:.4f} -> Best: {best_loss:.4f} ({final_impr:.1f}% improvement)')\n",
"print(f'Init PPL: {math.exp(init_loss):.1f} -> Best PPL: {math.exp(best_loss):.1f}')"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "generate"},
"outputs": [],
"source": [
"prompts = [\n",
" 'मैं आपको बताना चाहता हूँ कि',\n",
" 'भारत की राजधानी',\n",
" 'नमस्ते, आप कैसे हैं? मैं',\n",
"]\n",
"\n",
"for p in prompts:\n",
" inputs = tok(p, return_tensors='pt').to(model.device)\n",
" out = model.generate(\n",
" **inputs,\n",
" max_new_tokens=40,\n",
" temperature=0.3,\n",
" top_p=0.85,\n",
" repetition_penalty=1.25,\n",
" do_sample=True,\n",
" )\n",
" gen = tok.decode(out[0][inputs.input_ids.shape[-1]:],\n",
" skip_special_tokens=True)\n",
" print(f'Prompt: {p}')\n",
" print(f' -> {gen}')\n",
" print()"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {"id": "upload"},
"outputs": [],
"source": [
"from huggingface_hub import HfApi\n",
"\n",
"with torch.no_grad():\n",
" embed_w.data[new_ids] = learner.embed.data.to(embed_w.device)\n",
" lmhead_w.data[new_ids] = learner.lmhead.data.to(lmhead_w.device)\n",
"\n",
"model.save_pretrained('brahmi-trained')\n",
"tok.save_pretrained('brahmi-trained')\n",
"print('Saved locally')\n",
"\n",
"api = HfApi(token=HF_TOKEN)\n",
"api.upload_folder(\n",
" folder_path='brahmi-trained',\n",
" repo_id='eulogik/Bharat-Tiny-LLM-v2',\n",
" repo_type='model',\n",
" commit_message=f'continued_pretrain_round2: {final_impr:.1f}% val loss improvement',\n",
")\n",
"print('Uploaded to HF!')\n",
"print('https://huggingface.co/eulogik/Bharat-Tiny-LLM-v2')"
]
}
],
"metadata": {
"accelerator": "GPU",
"colab": {"provenance": []},
"kernelspec": {"display_name": "Python 3", "name": "python3"},
"language_info": {"name": "python"}
},
"nbformat": 4,
"nbformat_minor": 0
}