# ============================================================ # AYURVEDA GEMMA-3 4B TRAINING — FINAL v3 (Production) # Hardware: A100 MIG 3g.20gb (~19.5 GB VRAM) # Dataset: gemma4_ayurveda_unsloth_clean_WITH_REFS.jsonl (~971 K) # Model: unsloth/gemma-3-4b-it + QLoRA # Framework: Unsloth + TRL # ============================================================ # RESTART KERNEL BEFORE RUNNING CELL 1 # ============================================================ # ============================================================ # CELL 1 — ENVIRONMENT & CONFIG # ============================================================ import os import torch import gc # --- MUST BE SET BEFORE TORCH / UNSLOTH IMPORTS --- CACHE = "/nlsasfs/home/aikosh/prod-aikosh35/.cache/unsloth" os.makedirs(CACHE, exist_ok=True) os.makedirs(f"{CACHE}/torch", exist_ok=True) os.makedirs(f"{CACHE}/triton", exist_ok=True) os.environ["UNSLOTH_COMPILED_CACHE"] = CACHE os.environ["UNSLOTH_CACHE_DIR"] = CACHE os.environ["TORCHINDUCTOR_CACHE_DIR"] = f"{CACHE}/torch" os.environ["TRITON_CACHE_DIR"] = f"{CACHE}/triton" os.environ["XDG_CACHE_HOME"] = CACHE os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" # --- CONFIG --- MODEL_NAME = "unsloth/gemma-3-4b-it" DATASET_PATH = "/nlsasfs/home/aikosh/prod-aikosh35/gemma4_ayurveda_unsloth_clean_WITH_REFS.jsonl" OUTPUT_DIR = "/nlsasfs/home/aikosh/prod-aikosh35/gemma3_4b_ayurveda_output" MAX_SEQ_LENGTH = 1024 BATCH_SIZE = 1 GRAD_ACCUM = 32 LORA_R = 16 LORA_ALPHA = 16 LR = 2e-4 EPOCHS = 1 assert torch.cuda.is_available(), "CUDA GPU not found" print("=" * 60) print("GPU :", torch.cuda.get_device_name(0)) print("VRAM :", round(torch.cuda.get_device_properties(0).total_memory / 1024**3, 2), "GB") print("Model :", MODEL_NAME) print("Dataset:", DATASET_PATH) print("=" * 60) print("Cell 1 complete → Run Cell 2") # ============================================================ # CELL 2 — LOAD DATASET # ============================================================ from datasets import load_dataset dataset = load_dataset("json", data_files=DATASET_PATH, split="train") print(f"Loaded {len(dataset):,} conversations") sample = dataset[0] print("Roles :", [m["role"] for m in sample["messages"]]) print("\nUser preview:") print(sample["messages"][1]["content"][:200]) gc.collect() print("Cell 2 complete → Run Cell 3") # ============================================================ # CELL 3 — LOAD MODEL & TOKENIZER # ============================================================ from unsloth import FastLanguageModel print(f"Loading {MODEL_NAME} ...") try: model, tokenizer = FastLanguageModel.from_pretrained( model_name = MODEL_NAME, max_seq_length = MAX_SEQ_LENGTH, dtype = None, # Auto-detect (bf16 on A100) load_in_4bit = True, text_only = True, # Skip vision tower → ~2 GB VRAM saved ) except TypeError: # Fallback for older Unsloth builds that don't support text_only model, tokenizer = FastLanguageModel.from_pretrained( model_name = MODEL_NAME, max_seq_length = MAX_SEQ_LENGTH, dtype = None, load_in_4bit = True, ) print("VRAM used :", round(torch.cuda.memory_allocated() / 1024**3, 2), "GB") print("Chat template present:", tokenizer.chat_template is not None) gc.collect() torch.cuda.empty_cache() print("Cell 3 complete → Run Cell 4") # ============================================================ # CELL 4 — FORMAT DATASET + TOKEN DIAGNOSTICS + FILTER + SPLIT # ============================================================ def format_chat_template(examples): texts = [] for messages in examples["messages"]: text = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=False, # False = training mode ) texts.append(text) return {"text": texts} # Apply chat template dataset = dataset.map(format_chat_template, batched=True, remove_columns=["messages"]) print(f"Formatted {len(dataset):,} examples") # --- Token-length diagnostics (sample first 1 000 records) --- sample_tokens = [ len(tokenizer.encode(ex["text"])) for ex in dataset.select(range(min(1000, len(dataset)))) ] print("Max tokens :", max(sample_tokens)) print("Avg tokens :", round(sum(sample_tokens) / len(sample_tokens), 2)) # --- Filter overlong records (strict token limit) --- def keep_example(example): return len(tokenizer.encode(example["text"])) <= MAX_SEQ_LENGTH before = len(dataset) dataset = dataset.filter(keep_example) after = len(dataset) print(f"Removed {before - after:,} records > {MAX_SEQ_LENGTH} tokens") print(f"Remaining {after:,} records") # --- Train / Eval split (0.1 % eval ≈ 970 samples) --- split = dataset.train_test_split(test_size=0.001, seed=3407) train_dataset = split["train"] eval_dataset = split["test"] print("Train :", len(train_dataset)) print("Eval :", len(eval_dataset)) gc.collect() print("Cell 4 complete → Run Cell 5") # ============================================================ # CELL 5 — APPLY QLoRA ADAPTER (GUARDED) # ============================================================ from peft import PeftModel if isinstance(model, PeftModel): print("LoRA already attached — skipping.") else: print("Attaching LoRA adapter ...") model = FastLanguageModel.get_peft_model( model, r = LORA_R, lora_alpha = LORA_ALPHA, lora_dropout = 0, bias = "none", use_gradient_checkpointing = "unsloth", random_state = 3407, target_modules = [ "q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj", ], ) model.config.use_cache = False model.print_trainable_parameters() gc.collect() torch.cuda.empty_cache() print("VRAM after LoRA:", round(torch.cuda.memory_allocated() / 1024**3, 2), "GB") print("Cell 5 complete → Run Cell 6") # ============================================================ # CELL 6 — TRAINING CONFIGURATION # ============================================================ from trl import SFTConfig total_steps = (len(train_dataset) // (BATCH_SIZE * GRAD_ACCUM)) * EPOCHS print("Estimated steps:", total_steps) sft_args = SFTConfig( output_dir = OUTPUT_DIR, overwrite_output_dir = True, num_train_epochs = EPOCHS, per_device_train_batch_size = BATCH_SIZE, gradient_accumulation_steps = GRAD_ACCUM, learning_rate = LR, lr_scheduler_type = "cosine", warmup_steps = 100, optim = "adamw_8bit", weight_decay = 0.01, max_grad_norm = 0.3, bf16 = True, fp16 = False, logging_steps = 25, save_steps = 250, save_total_limit = 2, evaluation_strategy = "steps", eval_steps = 250, report_to = "none", dataset_text_field = "text", max_seq_length = MAX_SEQ_LENGTH, seed = 3407, ) print("Cell 6 complete → Run Cell 7") # ============================================================ # CELL 7 — INITIALIZE TRAINER # ============================================================ from trl import SFTTrainer trainer = SFTTrainer( model = model, train_dataset = train_dataset, eval_dataset = eval_dataset, processing_class= tokenizer, args = sft_args, packing = False, # Keep conversations separate ) print("Trainer ready") print("VRAM before training:", round(torch.cuda.memory_allocated() / 1024**3, 2), "GB") print("Cell 7 complete → Run Cell 8 (sanity) or Cell 9 (full training)") # ============================================================ # CELL 8 — OPTIONAL SANITY CHECK (100 records) # ============================================================ # Uncomment the block below to verify the pipeline on 100 samples # before launching the full ~971 K run. # # sanity_data = train_dataset.select(range(100)) # sanity_trainer = SFTTrainer( # model = model, # train_dataset = sanity_data, # processing_class = tokenizer, # args = sft_args, # packing = False, # ) # sanity_trainer.train() # print("Sanity check passed! Switch back to full dataset for real run.") # ============================================================ # CELL 9 — FULL TRAINING # ============================================================ gc.collect() torch.cuda.empty_cache() print("=" * 60) print("STARTING GEMMA-3 4B QLoRA TRAINING") print("=" * 60) print(f"Train set : {len(train_dataset):,}") print(f"Eval set : {len(eval_dataset):,}") print(f"Epochs : {EPOCHS}") print(f"Batch : {BATCH_SIZE} (accum: {GRAD_ACCUM}, effective: {BATCH_SIZE * GRAD_ACCUM})") print(f"Seq Len : {MAX_SEQ_LENGTH}") print(f"LoRA : r={LORA_R}, alpha={LORA_ALPHA}") print(f"LR : {LR}") print(f"VRAM : {torch.cuda.memory_allocated() / 1024**3:.1f} GB") print("=" * 60) # resume_from_checkpoint=True → starts from latest checkpoint if one exists in OUTPUT_DIR # safe to use even for a fresh run (starts from scratch if none found) trainer.train(resume_from_checkpoint=True) print("=" * 60) print("TRAINING COMPLETE") print("=" * 60) print("Cell 9 complete → Run Cell 10") # ============================================================ # CELL 10 — SAVE ADAPTER + OPTIONAL MERGED 16-BIT MODEL # ============================================================ os.makedirs(OUTPUT_DIR, exist_ok=True) # --- LoRA adapter --- adapter_path = os.path.join(OUTPUT_DIR, "adapter") model.save_pretrained(adapter_path) tokenizer.save_pretrained(adapter_path) print("Adapter saved to:", adapter_path) # --- Optional: merged 16-bit weights (easier for downstream inference) --- try: merged_path = os.path.join(OUTPUT_DIR, "merged_16bit") model.save_pretrained_merged( merged_path, tokenizer, save_method="merged_16bit", ) print("Merged 16-bit model saved to:", merged_path) except Exception as e: print("Merged save skipped:", e) print("Cell 10 complete → Run Cell 11 for inference test") # ============================================================ # CELL 11 — INFERENCE TEST # ============================================================ from unsloth import FastLanguageModel FastLanguageModel.for_inference(model) messages = [ { "role": "user", "content": ( "What does Charaka say about the qualities of a good physician " "in Sutrasthana?" ), } ] inputs = tokenizer.apply_chat_template( messages, tokenize=True, return_tensors="pt", add_generation_prompt=True, # True = inference mode ).to("cuda") print("Generating ...") outputs = model.generate( inputs, max_new_tokens = 512, temperature = 0.7, top_p = 0.9, do_sample = True, ) response = tokenizer.decode(outputs[0], skip_special_tokens=True) print("\n" + "=" * 60) print("RESPONSE") print("=" * 60) print(response) print("=" * 60)