Add Kaggle training notebook (live run)
Browse files
training/Kaggle_GRPO_Training_LIVE.ipynb
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"accelerator":"GPU","colab":{"gpuType":"T4","provenance":[]},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"datasetVersion","sourceId":1448850,"datasetId":849299,"databundleVersionId":1482399,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":1157702,"datasetId":654897,"databundleVersionId":1188575},{"sourceType":"datasetVersion","sourceId":7639866,"datasetId":841565,"databundleVersionId":7736362},{"sourceType":"datasetVersion","sourceId":468668,"datasetId":215919,"databundleVersionId":484566,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🏥 TRIAGE GRPO Training — Colab Edition\n\n**Meta PyTorch OpenEnv Hackathon** | 14 datasets | 9 reward verifiers | GRPO\n\n| Item | Value |\n|---|---|\n| Model | Qwen2.5-7B (4-bit NF4) |\n| Datasets | 14 sources (7 HF + 6 Kaggle + 1 base) |\n| Method | GRPO with 9 reward verifiers |\n| Hardware | Colab T4 / Kaggle P100 |","metadata":{}},{"cell_type":"code","source":"!pip install -q \"transformers>=4.45\" \"trl>=0.12\" \"peft>=0.13\" \"bitsandbytes>=0.46\" \"datasets>=3.0\" \"accelerate>=1.0\" huggingface_hub kagglehub pyarrow","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T06:34:11.011542Z","iopub.execute_input":"2026-04-26T06:34:11.011946Z","iopub.status.idle":"2026-04-26T06:34:18.645358Z","shell.execute_reply.started":"2026-04-26T06:34:11.011899Z","shell.execute_reply":"2026-04-26T06:34:18.644379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json,re,random,logging,time,os,gc,torch,inspect\nfrom pathlib import Path\nlogging.basicConfig(level=logging.INFO,format='%(asctime)s %(levelname)s %(message)s')\nCFG={'model':'Qwen/Qwen2.5-7B','max_seq_length':512,'lora_r':16,'lora_alpha':32,'lora_dropout':0,\n 'lora_targets':['q_proj','k_proj','v_proj','o_proj','gate_proj','up_proj','down_proj'],\n 'num_generations':4,'max_completion_length':200,'temperature':0.9,\n 'epochs':1,'batch_size':1,'grad_accum':4,'lr':5e-5,\n 'logging_steps':5,'save_steps':50,'output_dir':'./grpo_output','base_dataset':'balarajr/triage-grpo'}\nHF_TOKEN=os.environ.get('HF_TOKEN','')\nUSE_BF16=torch.cuda.is_available() and torch.cuda.is_bf16_supported()\nCOMPUTE_DTYPE=torch.bfloat16 if USE_BF16 else torch.float16\nprint(f'GPU: {torch.cuda.get_device_name(0) if torch.cuda.is_available() else \"CPU\"}')\nprint(f'Dtype: {\"bf16\" if USE_BF16 else \"fp16\"}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T06:34:18.647414Z","iopub.execute_input":"2026-04-26T06:34:18.647767Z","iopub.status.idle":"2026-04-26T06:34:23.758664Z","shell.execute_reply.started":"2026-04-26T06:34:18.647738Z","shell.execute_reply":"2026-04-26T06:34:23.757963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 9 Reward Verifiers\n_VA=frozenset({'TRIAGE_PATIENT','ASSIGN_TREATMENT','TRANSFER_TO_ICU','TRANSFER_TO_WARD','ACTIVATE_OVERFLOW','ORDER_MEDICATION','FLAG_POLICY_VIOLATION','OVERRIDE_DECISION','UPDATE_EHR','REQUEST_STAFF','VERIFY_INSURANCE'})\n_RK={'action_type','target_id','priority','reasoning'}\n_EV=[r'P-\\d{2,3}',r'patient\\s+\\d+',r'\\d+%',r'\\d+/\\d+',r'BP\\s*\\d+',r'HR\\s*\\d+',r'ICU\\s+at\\s+\\d+',r'beds?\\s+\\d+',r'age\\s+\\d+',r'critical|immediate|urgent|stable']\n_FL=['i need more information',\"i'm not sure\",'let me think','i cannot determine',\"i don't know\",'more data needed']\n_FB=[r'\\bimport\\s+os\\b',r'\\bimport\\s+sys\\b',r'\\bexec\\s*\\(',r'\\beval\\s*\\(',r'\\breward\\s*[:=]\\s*1\\.0\\b']\n\ndef _xj(text):\n text=text.strip()\n try: return json.loads(text)\n except: pass\n m=re.search(r'```(?:json)?\\s*(\\{.*?\\})\\s*```',text,re.DOTALL)\n if m:\n try: return json.loads(m.group(1))\n except: pass\n m=re.search(r'\\{[^{}]*\\}',text,re.DOTALL)\n if m:\n try: return json.loads(m.group(0))\n except: pass\n return None\n\ndef _sp(prompt):\n s={'alive_count':20,'deceased_count':0,'critical_count':0,'icu_occupancy':0.5,'violations_injected':0,'violations_caught':0,'survival_rate':1.0}\n m=re.search(r'ICU OCCUPANCY:\\s*(\\d+)%',prompt)\n if m: s['icu_occupancy']=int(m.group(1))/100.0\n m=re.search(r'CRITICAL PATIENTS\\s*\\((\\d+)',prompt)\n if m: s['critical_count']=int(m.group(1))\n m=re.search(r'VIOLATIONS INJECTED:\\s*(\\d+)\\s*\\|\\s*CAUGHT:\\s*(\\d+)',prompt)\n if m: s['violations_injected'],s['violations_caught']=int(m.group(1)),int(m.group(2))\n m=re.search(r'SURVIVAL RATE:\\s*(\\d+\\.?\\d*)%',prompt)\n if m: s['survival_rate']=float(m.group(1))/100.0\n t=20;s['alive_count']=int(s['survival_rate']*t);s['deceased_count']=t-s['alive_count']\n return s\n\ndef reward_format_compliance(completions,**kw):\n R=[]\n for c in completions:\n p=_xj(c)\n if not p or not _RK.issubset(p.keys()): R.append(0.0);continue\n a=str(p.get('action_type','')).upper()\n if a not in _VA: R.append(0.0);continue\n try: int(p['target_id'])\n except: R.append(0.0);continue\n try:\n pr=int(p['priority'])\n if not 1<=pr<=10: R.append(0.0);continue\n except: R.append(0.0);continue\n if len(str(p.get('reasoning','')).strip())<10: R.append(0.0);continue\n R.append(1.0)\n return R\n\ndef reward_patient_survival(completions,**kw):\n P=kw.get('prompts',kw.get('prompt',['']));R=[]\n for i,c in enumerate(completions):\n s=_sp(P[i] if i<len(P) else '');t=s['alive_count']+s['deceased_count']\n R.append(s['alive_count']/t if t>0 else 1.0)\n return R\n\ndef reward_icu_efficiency(completions,**kw):\n P=kw.get('prompts',kw.get('prompt',['']));R=[]\n for i,c in enumerate(completions):\n o=_sp(P[i] if i<len(P) else '')['icu_occupancy']\n R.append(1.0 if o<=0.85 else (1.0-(o-0.85)*5.0 if o<=0.95 else max(0.0,0.5-(o-0.95)*10.0)))\n return R\n\ndef reward_violation_detection(completions,**kw):\n P=kw.get('prompts',kw.get('prompt',['']));R=[]\n for i,c in enumerate(completions):\n s=_sp(P[i] if i<len(P) else '');inj,ct=s['violations_injected'],s['violations_caught']\n R.append(min(1.0,ct/max(inj,1)) if inj>0 else 1.0)\n return R\n\ndef reward_reasoning_quality(completions,**kw):\n R=[]\n for c in completions:\n p=_xj(c)\n if not p: R.append(0.0);continue\n r=str(p.get('reasoning',''))\n if len(r)<20: R.append(0.1);continue\n ev=sum(1 for pat in _EV if re.search(pat,r,re.I))\n if any(f in r.lower() for f in _FL): R.append(0.1);continue\n R.append(min(1.0,0.3+min(0.7,ev*0.15)))\n return R\n\ndef reward_response_speed(completions,**kw):\n return [1.0 if len(c)<=400 else (1.0-(len(c)-400)*0.001 if len(c)<=800 else max(0.2,0.6-(len(c)-800)*0.0005)) for c in completions]\n\ndef reward_no_hallucination(completions,**kw):\n P=kw.get('prompts',kw.get('prompt',['']));R=[]\n for i,c in enumerate(completions):\n p=_xj(c)\n if not p: R.append(0.5);continue\n mn={int(m.group(1)) for m in re.finditer(r'P-(\\d{2,3})',str(p.get('reasoning','')),re.I)}\n if not mn: R.append(1.0);continue\n vl={int(m.group(1)) for m in re.finditer(r'P-(\\d{2,3})',P[i] if i<len(P) else '')}\n R.append(0.0 if mn-vl else 1.0)\n return R\n\ndef reward_action_alignment(completions,**kw):\n P=kw.get('prompts',kw.get('prompt',['']));R=[]\n for i,c in enumerate(completions):\n p=_xj(c)\n if not p: R.append(0.0);continue\n s=_sp(P[i] if i<len(P) else '');a=str(p.get('action_type','')).upper()\n o,cr=s['icu_occupancy'],s['critical_count'];v=s['violations_injected']-s['violations_caught']\n sm={'TRIAGE_PATIENT':1.0 if cr>0 else 0.5,'TRANSFER_TO_ICU':1.0 if o<0.9 and cr>0 else 0.3,'ACTIVATE_OVERFLOW':1.0 if o>=0.85 else 0.2,'FLAG_POLICY_VIOLATION':1.0 if v>0 else 0.4,'ORDER_MEDICATION':0.8 if cr>0 else 0.5,'ASSIGN_TREATMENT':0.9 if cr>0 else 0.5}\n R.append(sm.get(a,0.5))\n return R\n\ndef reward_sandbox_safety(completions,**kw):\n R=[]\n for c in completions:\n safe=all(not re.search(pat,c,re.I) for pat in _FB) and len(c)<=3000\n w=c.split()\n if len(w)>20 and len(set(w))/len(w)<0.2: safe=False\n R.append(1.0 if safe else 0.0)\n return R\n\nREWARD_FUNCS=[reward_format_compliance,reward_patient_survival,reward_icu_efficiency,reward_violation_detection,reward_reasoning_quality,reward_response_speed,reward_no_hallucination,reward_action_alignment,reward_sandbox_safety]\nprint(f'Loaded {len(REWARD_FUNCS)} reward verifiers')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T06:34:23.759642Z","iopub.execute_input":"2026-04-26T06:34:23.760078Z","iopub.status.idle":"2026-04-26T06:34:23.786434Z","shell.execute_reply.started":"2026-04-26T06:34:23.760051Z","shell.execute_reply":"2026-04-26T06:34:23.785716Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4: Load Model + LoRA\nfrom transformers import AutoModelForCausalLM,AutoTokenizer,BitsAndBytesConfig\nfrom peft import LoraConfig,get_peft_model,prepare_model_for_kbit_training\nbnb=BitsAndBytesConfig(load_in_4bit=True,bnb_4bit_quant_type='nf4',bnb_4bit_compute_dtype=COMPUTE_DTYPE,bnb_4bit_use_double_quant=True)\nprint(f'Loading {CFG[\"model\"]} ...')\nmodel=AutoModelForCausalLM.from_pretrained(CFG['model'],quantization_config=bnb,device_map='auto',torch_dtype=COMPUTE_DTYPE,trust_remote_code=True,token=HF_TOKEN or None)\ntokenizer=AutoTokenizer.from_pretrained(CFG['model'],trust_remote_code=True,token=HF_TOKEN or None)\nif tokenizer.pad_token is None:\n tokenizer.pad_token=tokenizer.eos_token;model.config.pad_token_id=model.config.eos_token_id\nmodel=prepare_model_for_kbit_training(model)\nlora=LoraConfig(r=CFG['lora_r'],lora_alpha=CFG['lora_alpha'],lora_dropout=CFG['lora_dropout'],target_modules=CFG['lora_targets'],bias='none',task_type='CAUSAL_LM')\nmodel=get_peft_model(model,lora)\nmodel.print_trainable_parameters()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T06:34:23.787466Z","iopub.execute_input":"2026-04-26T06:34:23.787886Z","iopub.status.idle":"2026-04-26T06:36:13.027303Z","shell.execute_reply.started":"2026-04-26T06:34:23.787851Z","shell.execute_reply":"2026-04-26T06:36:13.026649Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 5: Multi-Source Dataset Pipeline (14 Sources)\n7 HuggingFace + 6 Kaggle + 1 base triage dataset","metadata":{}},{"cell_type":"code","source":"import pandas as pd,kagglehub\nfrom datasets import load_dataset,Dataset\nrng=random.Random(42)\nAG=['ER_TRIAGE','ICU_MANAGEMENT','PHARMACY','CMO_OVERSIGHT','HR_ROSTERING','IT_SYSTEMS']\nCR=['MASS_CASUALTY','OUTBREAK','EQUIPMENT_FAILURE','STAFF_SHORTAGE']\nJS='Respond with ONLY valid JSON:\\n{\\n \"action_type\": \"<action>\",\\n \"target_id\": <int>,\\n \"priority\": <1-10>,\\n \"reasoning\": \"<cite data>\"\\n}'\n\ndef mh():\n a=rng.choice(AG);c=rng.choice(CR);icu=rng.randint(30,98);cr=rng.randint(1,10)\n return (f'You are the {a} agent in a hospital crisis simulation.\\n\\nCRISIS: {c}\\nSTEP: {rng.randint(0,19)}/20\\n'\n f'ICU OCCUPANCY: {icu}% ({icu*20//100}/20 beds)\\nCRITICAL PATIENTS ({cr} total):\\n'\n f' P-{rng.randint(1,99):03d}: CRITICAL -- BP {rng.randint(55,95)}/{rng.randint(25,65)}, HR {rng.randint(95,160)}\\n'\n f'VIOLATIONS INJECTED: {rng.randint(0,5)} | CAUGHT: {rng.randint(0,3)}\\nSURVIVAL RATE: {rng.uniform(82,100):.1f}%\\n\\n')\n\nall_prompts=[]\n\n# [0] Base triage dataset\nprint('[0/14] balarajr/triage-grpo')\ntry:\n ds=load_dataset(CFG['base_dataset'],split='train',token=HF_TOKEN or None);all_prompts.extend(list(ds['prompt']))\n print(f' \\u2714 {len(ds)}')\nexcept Exception as e: print(f' \\u2718 {e}')\n\n# HuggingFace datasets\nHF_SOURCES=[\n ('FreedomIntelligence/medical-o1-reasoning-SFT','data/train-00000-of-00001.parquet','CLINICAL',500),\n ('bigbio/med_qa','med_qa_en_bigbio_qa/train-00000-of-00001.parquet','MEDQA',500),\n ('sdiazlor/medical-reasoning-dataset','data/train-00000-of-00001.parquet','REASONING',500),\n ('Anthropic/hh-rlhf','data/harmless-base/train-00000-of-00001.parquet','SAFETY',300),\n ('PKU-Alignment/PKU-SafeRLHF','data/train-00000-of-00001.parquet','ALIGNMENT',300),\n ('lavita/ChatDoctor-iCliniq','data/train-00000-of-00001.parquet','CONSULT',500),\n]\nfor i,(slug,path,tag,n) in enumerate(HF_SOURCES,1):\n print(f'[{i}/14] {slug}')\n try:\n df=pd.read_parquet(f'hf://datasets/{slug}/{path}')\n for _,row in df.sample(min(n,len(df)),random_state=42).iterrows():\n txt=str(row.get('question',row.get('input',row.get('instruction',row.get('chosen',row.get('prompt',''))))))[:300]\n all_prompts.append(mh()+f'{tag}: {txt}\\n\\n{JS}')\n print(f' \\u2714 {min(n,len(df))}/{len(df)}')\n except Exception as e: print(f' \\u2718 {e}')\n\n# [7] medical flashcards (JSON format)\nprint('[7/14] medalpaca/medical_meadow_medical_flashcards')\ntry:\n df=pd.read_json('hf://datasets/medalpaca/medical_meadow_medical_flashcards/medical_meadow_medical_flashcards.json')\n for _,row in df.sample(min(500,len(df)),random_state=42).iterrows():\n txt=str(row.get('input',row.get('instruction','')))[:300]\n all_prompts.append(mh()+f'FLASHCARD: {txt}\\n\\n{JS}')\n print(f' \\u2714 500/{len(df)}')\nexcept Exception as e: print(f' \\u2718 {e}')\n\n# Kaggle datasets\nKG_SOURCES=[\n ('thedevastator/medical-q-a-structured','STRUCT_QA',500),\n ('nehaprabhavalkar/av-healthcare-analytics-ii','ANALYTICS',500),\n ('jpmiller/layoutlm','NLP_REC',300),\n ('thedevastator/usmle-medical-licensing-examination','USMLE',500),\n ('kaushil268/disease-prediction-using-machine-learning','DISEASE',500),\n ('maalona/hospital-triage-and-patient-history-data','TRIAGE_HIST',500),\n]\nfor i,(slug,tag,n) in enumerate(KG_SOURCES,8):\n print(f'[{i}/14] Kaggle: {slug}')\n try:\n p=kagglehub.dataset_download(slug)\n csvs=[f for f in os.listdir(p) if f.endswith('.csv')]\n if not csvs: print(' \\u2718 no CSV');continue\n df=pd.read_csv(os.path.join(p,csvs[0]),nrows=n*3)\n for _,row in df.sample(min(n,len(df)),random_state=42).iterrows():\n cols=' | '.join(f'{c}: {row[c]}' for c in df.columns[:6])[:300]\n all_prompts.append(mh()+f'{tag}: {cols}\\n\\n{JS}')\n print(f' \\u2714 {min(n,len(df))}/{len(df)}')\n except Exception as e: print(f' \\u2718 {e}')\n\nrandom.shuffle(all_prompts)\ndataset=Dataset.from_dict({'prompt':all_prompts})\nprint(f'\\n\\u2501'*60)\nprint(f'\\u2705 TOTAL: {len(dataset)} prompts from 14 sources')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T06:36:13.028922Z","iopub.execute_input":"2026-04-26T06:36:13.029928Z","iopub.status.idle":"2026-04-26T06:39:55.079405Z","shell.execute_reply.started":"2026-04-26T06:36:13.0299Z","shell.execute_reply":"2026-04-26T06:39:55.078526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 6: GRPO Training\nfrom trl import GRPOTrainer,GRPOConfig\ngk=dict(output_dir=CFG['output_dir'],num_train_epochs=CFG['epochs'],per_device_train_batch_size=CFG['batch_size'],\n gradient_accumulation_steps=CFG['grad_accum'],learning_rate=CFG['lr'],max_completion_length=CFG['max_completion_length'],\n num_generations=CFG['num_generations'],temperature=CFG['temperature'],logging_steps=CFG['logging_steps'],\n save_steps=CFG['save_steps'],save_total_limit=2,report_to='none',bf16=USE_BF16,fp16=not USE_BF16,seed=42)\nsig=inspect.signature(GRPOConfig)\nif not any(p.kind==inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()):\n gk={k:v for k,v in gk.items() if k in set(sig.parameters)}\nargs=GRPOConfig(**gk)\ntk=dict(model=model,reward_funcs=REWARD_FUNCS,args=args,train_dataset=dataset)\ntp=set(inspect.signature(GRPOTrainer.__init__).parameters)\nif 'processing_class' in tp: tk['processing_class']=tokenizer\nelif 'tokenizer' in tp: tk['tokenizer']=tokenizer\ntrainer=GRPOTrainer(**tk)\nprint(f'Starting GRPO ({len(dataset)} prompts, {CFG[\"epochs\"]} epochs)...')\nresult=trainer.train()\nprint(f'Done! Loss={result.training_loss:.4f} Steps={result.global_step}')\ntrainer.save_model(CFG['output_dir']);tokenizer.save_pretrained(CFG['output_dir'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T06:39:55.080603Z","iopub.execute_input":"2026-04-26T06:39:55.081292Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 7: Quick Evaluation\nfrom peft import PeftModel\ndel model;gc.collect();torch.cuda.empty_cache()\nbase=AutoModelForCausalLM.from_pretrained(CFG['model'],quantization_config=bnb,device_map='auto',torch_dtype=COMPUTE_DTYPE,trust_remote_code=True,token=HF_TOKEN or None)\nmodel=PeftModel.from_pretrained(base,CFG['output_dir'])\ntok=AutoTokenizer.from_pretrained(CFG['output_dir'],trust_remote_code=True)\ntp=('You are the ER_TRIAGE agent in a hospital crisis simulation.\\n\\nCRISIS: MASS_CASUALTY\\nSTEP: 5/20\\n'\n 'ICU OCCUPANCY: 85% (17/20 beds)\\nCRITICAL PATIENTS (3 total):\\n P-042: CRITICAL -- BP 72/40, HR 140\\n'\n ' P-019: CRITICAL -- BP 65/35, HR 155\\n P-067: CRITICAL -- BP 80/50, HR 120\\n'\n 'VIOLATIONS INJECTED: 2 | CAUGHT: 1\\nSURVIVAL RATE: 90.0%\\n\\n'\n 'Respond with ONLY valid JSON:\\n{\\n \"action_type\": \"<action>\",\\n \"target_id\": <int>,\\n \"priority\": <1-10>,\\n \"reasoning\": \"<cite data>\"\\n}')\ninputs=tok(tp,return_tensors='pt').to(model.device)\nwith torch.no_grad():\n out=model.generate(**inputs,max_new_tokens=200,temperature=0.7,do_sample=True)\nresp=tok.decode(out[0][inputs['input_ids'].shape[1]:],skip_special_tokens=True)\nprint('='*60);print('EVAL RESPONSE:');print(resp);print('='*60)\nparsed=_xj(resp)\nif parsed:\n scores={fn.__name__:fn([resp],prompts=[tp])[0] for fn in REWARD_FUNCS}\n total=sum(scores.values())/len(scores)*100\n print(f'\\nReward Scores: {json.dumps(scores,indent=2)}')\n print(f'Overall: {total:.0f}/100')\nelse: print('WARNING: Could not parse JSON')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}
|