Files
Bharat-Tiny-LLM-v2/brahmi_continued_pretrain.ipynb

395 lines
14 KiB
Plaintext
Raw Normal View History

{
"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
}