{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from pathlib import Path\nimport os,sys,shutil,subprocess,re,yaml,time\nimport pandas as pd\n\nT0=time.time()\ndef log(x): print(f\"[{(time.time()-T0)/60:6.1f}m] {x}\",flush=True)\n\nROOT=Path(\"/kaggle/working/rxrx1\")\nCOMP1=Path(\"/kaggle/input/competitions/recursion-cellular-image-classification\")\nCOMP2=Path(\"/kaggle/input/recursion-cellular-image-classification\")\nCOMP=COMP1 if COMP1.exists() else COMP2\nDATA=Path(\"/kaggle/input/rxrx1-epoch15-kaggle-compare\")\nMETA=DATA/\"rxrx1-epoch15-kaggle\"\n\nCKPT=DATA/\"epoch_03_trainacc_0.5254.pt\"\nTREAT=META/\"full_all_treatment.csv\"\nRAW_NAME=\"pseudo_03_trainnorm_raw.csv\"\nLSA_NAME=\"pseudo_03_trainnorm_lsa.csv\"\n\nTRAIN_MEAN=9.640125374750943\nTRAIN_STD=13.499834120240363\n\nassert COMP.exists(),COMP\nassert CKPT.exists(),CKPT\nassert TREAT.exists(),TREAT\nlog(f\"START | ckpt={CKPT.name}\")\n\n# ============================================================\n# 1. CLONE\n# ============================================================\n\nos.chdir(\"/kaggle/working\")\nif ROOT.exists(): shutil.rmtree(ROOT)\n\nlog(\"cloning repo\")\nsubprocess.run([\n    \"git\",\"clone\",\"-b\",\"feat/adaptive-final-training\",\n    \"https://github.com/inori6/rxrx1.git\",str(ROOT)\n],check=True)\n\nos.chdir(ROOT)\nos.environ[\"PYTHONPATH\"]=str(ROOT/\"src\")\nos.environ[\"PYTHONUNBUFFERED\"]=\"1\"\ncommit=subprocess.check_output([\"git\",\"rev-parse\",\"--short\",\"HEAD\"],text=True).strip()\nlog(f\"clone done | commit={commit}\")\n\n# ============================================================\n# 2. LINK TEST DATA\n# ============================================================\n\nraw=ROOT/\"data/raw\"\nraw.mkdir(parents=True,exist_ok=True)\n\nfor fn in [\"test.csv\",\"sample_submission.csv\"]:\n    dst=raw/fn\n    if dst.exists() or dst.is_symlink(): dst.unlink()\n    dst.symlink_to(COMP/fn)\n\ndst=raw/\"test\"\nif dst.exists() or dst.is_symlink():\n    dst.unlink() if dst.is_symlink() else shutil.rmtree(dst)\ndst.symlink_to(COMP/\"test\",target_is_directory=True)\nlog(\"test data linked\")\n\n# ============================================================\n# 3. MANIFEST\n# ============================================================\n\nsplit=ROOT/\"data/processed/splits\"\nsplit.mkdir(parents=True,exist_ok=True)\nshutil.copy2(TREAT,split/\"full_all_treatment.csv\")\nlog(\"manifest ready\")\n\n# ============================================================\n# 4. PATCH PREDICTOR\n# ============================================================\n\npred=ROOT/\"scripts/predict_test_submit.py\"\ntxt=pred.read_text()\n\nif \"import os\\n\" not in txt:\n    txt=txt.replace(\"import argparse\\n\",\"import argparse\\nimport os\\n\",1)\n\ntxt,k=re.subn(\n    r\"def expand_complete_test_sites\\(test, test_root\\):.*?\\n\\ndef build_test_stats_normalizer\",\n    '''def expand_complete_test_sites(test, test_root):\n    return test.merge(pd.DataFrame({\"site\":[1,2]}),how=\"cross\").reset_index(drop=True)\n\ndef build_test_stats_normalizer''',\n    txt,flags=re.S,count=1\n)\nassert k==1,\"site scan patch failed\"\n\ntxt,k=re.subn(\n    r\"def build_inference_normalizer\\(.*?\\n\\ndef main\\(\\):\",\n    '''def build_inference_normalizer(config,test_manifest,label_to_index,root,norm_stats_source):\n    from rxrx1.data.normalization import NormStats,NormStatsStore\n\n    mean=float(os.environ[\"RXRX1_FIXED_NORM_MEAN\"])\n    std=float(os.environ[\"RXRX1_FIXED_NORM_STD\"])\n\n    stats=NormStatsStore()\n    stats.add(\n        \"global\",\n        NormStats(\n            mean=torch.tensor([[[mean]]],dtype=torch.float32),\n            std=torch.tensor([[[std]]],dtype=torch.float32),\n            count=1,\n            element_count=1,\n        ),\n    )\n\n    section=config.get(\"normalization\") or {}\n    normalizer=ReferenceZScoreNormalizer(\n        stats=stats,\n        grouping=\"global\",\n        eps=float(section.get(\"eps\",1e-6)),\n        missing_group=str(section.get(\"missing_group\",\"error\")).lower(),\n    )\n    normalizer.position=\"before_resize\"\n    print(f\"FIXED {norm_stats_source.upper()} NORM | mean={mean} | std={std}\",flush=True)\n    return normalizer\n\ndef main():''',\n    txt,flags=re.S,count=1\n)\nassert k==1,\"normalization patch failed\"\n\ntxt,k=re.subn(\n    r'    well_predictions = pd\\.DataFrame\\(well_rows\\)\\n\\n    if lsa_enabled:',\n    '''    well_predictions = pd.DataFrame(well_rows)\n\n    raw_internal_predictions={}\n    for row in well_rows:\n        class_index=int(row[\"logits\"].argmax())\n        raw_internal_predictions[str(row[\"id_code\"])]=index_to_label[class_index]\n\n    raw_predictions={\n        id_code:int(sirna_to_official[label])\n        for id_code,label in raw_internal_predictions.items()\n    }\n\n    raw_submission=sample[[\"id_code\"]].copy()\n    raw_submission[\"sirna\"]=raw_submission[\"id_code\"].astype(str).map(raw_predictions)\n\n    if raw_submission[\"sirna\"].isna().any():\n        missing_ids=raw_submission.loc[\n            raw_submission[\"sirna\"].isna(),\"id_code\"\n        ].astype(str).tolist()[:10]\n        raise RuntimeError(f\"Missing RAW predictions: {missing_ids}\")\n\n    raw_submission[\"sirna\"]=raw_submission[\"sirna\"].astype(int)\n    raw_output_path=Path(os.environ[\"RXRX1_RAW_OUTPUT\"])\n    raw_output_path.parent.mkdir(parents=True,exist_ok=True)\n    raw_submission.to_csv(raw_output_path,index=False)\n    print(f\"RAW submission saved: {raw_output_path}\",flush=True)\n\n    if lsa_enabled:''',\n    txt,count=1\n)\nassert k==1,\"raw output patch failed\"\n\npred.write_text(txt)\nlog(\"predictor patched\")\n\n# ============================================================\n# 5. CONFIG\n# ============================================================\n\ncfg_path=ROOT/\"configs/final_selection/02_aggressive.yaml\"\ncfg=yaml.safe_load(cfg_path.read_text())\n\ncfg[\"model\"][\"pretrained\"]=False\ncfg[\"data\"][\"batch_size\"]=32\n\ncfg.setdefault(\"inference\",{})\ncfg[\"inference\"].setdefault(\"lsa\",{})\ncfg[\"inference\"][\"lsa\"][\"enabled\"]=True\ncfg[\"inference\"].setdefault(\"tta\",{})\ncfg[\"inference\"][\"tta\"][\"enabled\"]=False\n\ncfg_path.write_text(yaml.safe_dump(cfg,sort_keys=False))\nlog(\"config ready | BS=32 | LSA=ON | TTA=OFF\")\n\n# ============================================================\n# 6. INFERENCE\n# ============================================================\n\nout=ROOT/\"outputs/submissions\"\nout.mkdir(parents=True,exist_ok=True)\n\nenv=os.environ.copy()\nenv[\"PYTHONUNBUFFERED\"]=\"1\"\nenv[\"RXRX1_FIXED_NORM_MEAN\"]=repr(TRAIN_MEAN)\nenv[\"RXRX1_FIXED_NORM_STD\"]=repr(TRAIN_STD)\nenv[\"RXRX1_RAW_OUTPUT\"]=str(out/RAW_NAME)\n\ncmd=[\n    sys.executable,\"-u\",\"scripts/predict_test_submit.py\",\n    \"--config\",str(cfg_path),\n    \"--checkpoint\",str(CKPT),\n    \"--output\",str(out/LSA_NAME),\n    \"--norm-stats-source\",\"train\",\n]\n\nlog(\"INFERENCE START\")\np=subprocess.Popen(cmd,cwd=ROOT,env=env)\n\nwhile True:\n    rc=p.poll()\n    if rc is not None: break\n    time.sleep(60)\n    log(\"inference still running...\")\n\nif rc!=0:\n    raise subprocess.CalledProcessError(rc,cmd)\n\nlog(\"INFERENCE DONE\")\n\n# ============================================================\n# 7. CHECK + EXPORT\n# ============================================================\n\ndfs={}\nfor kind,fn in [(\"RAW\",RAW_NAME),(\"LSA\",LSA_NAME)]:\n    path=out/fn\n    assert path.exists(),path\n    df=pd.read_csv(path)\n    assert len(df)==df.id_code.nunique()\n    assert df.sirna.between(0,1107).all()\n    shutil.copy2(path,Path(\"/kaggle/working\")/fn)\n    dfs[kind]=df\n    log(f\"{kind} | rows={len(df)} | classes={df.sirna.nunique()}\")\n\ncmp=dfs[\"RAW\"].merge(dfs[\"LSA\"],on=\"id_code\",suffixes=(\"_raw\",\"_lsa\"))\nchanged=(cmp.sirna_raw!=cmp.sirna_lsa).sum()\n\nlog(f\"ALL DONE | LSA changed={changed}/{len(cmp)} ({changed/len(cmp):.2%})\")\nprint(f\"RAW: /kaggle/working/{RAW_NAME}\",flush=True)\nprint(f\"LSA: /kaggle/working/{LSA_NAME}\",flush=True)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}