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