{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import gc\nimport json\nimport os\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nimport sys\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\n\n# ---- user-configurable paths ----\nDEFAULT_TEST_CSV = \"/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv\"\nDEFAULT_OUTPUT = \"/kaggle/working/submission.csv\"\nDEFAULT_CODE_DIR = \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1\"# \"/kaggle/input/protenix-v1/Protenix-v1\"\nDEFAULT_ROOT_DIR = \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1\"# \"/kaggle/input/protenix-v1/Protenix-v1\"\nMODEL_NAME = \"protenix_base_default_v1.0.0\"\nN_SAMPLE = 5\nSEED = 42\n\n\nMAX_SEQ_LEN = int(os.environ.get(\"MAX_SEQ_LEN\", \"512\"))\nCHUNK_OVERLAP = int(os.environ.get(\"CHUNK_OVERLAP\", \"64\"))\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\n\n#如果 TEST_TARGETS 为空，就按全量运行。\n\n#TEST_TARGETS = os.environ.get(\"TEST_TARGETS\", \"9MME,9ZCC\").strip()\nTEST_TARGETS = os.environ.get(\"TEST_TARGETS\", \"8ZNQ\").strip()\n\n\ndef parse_bool(value: str, default: bool = False) -> str:\n    normalized = str(value).strip().lower()\n    if normalized in {\"1\", \"true\", \"t\", \"yes\", \"y\", \"on\"}:\n        return \"true\"\n    if normalized in {\"0\", \"false\", \"f\", \"no\", \"n\", \"off\"}:\n        return \"false\"\n    return \"true\" if default else \"false\"\n\n\nUSE_MSA = parse_bool(os.environ.get(\"USE_MSA\", \"false\"))\nUSE_TEMPLATE = parse_bool(os.environ.get(\"USE_TEMPLATE\", \"false\"))\nUSE_RNA_MSA = parse_bool(os.environ.get(\"USE_RNA_MSA\", \"false\"))\n\nTRIANGLE_ATTENTION = os.environ.get(\"TRIANGLE_ATTENTION\", \"torch\").strip().lower()\nTRIANGLE_MULTIPLICATIVE = os.environ.get(\"TRIANGLE_MULTIPLICATIVE\", \"torch\").strip().lower()\n\ndef seed_everything(seed: int) -> None:\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.enabled = True\n    torch.use_deterministic_algorithms(True)\n\n\ndef resolve_paths() -> tuple[str, str, str]:\n    test_csv = os.environ.get(\"TEST_CSV\", DEFAULT_TEST_CSV)\n    output_csv = os.environ.get(\"SUBMISSION_CSV\", DEFAULT_OUTPUT)\n    code_dir = os.environ.get(\"PROTENIX_CODE_DIR\", DEFAULT_CODE_DIR)\n    root_dir = os.environ.get(\"PROTENIX_ROOT_DIR\", DEFAULT_ROOT_DIR)\n    return test_csv, output_csv, code_dir, root_dir\n\n\ndef ensure_required_files(root_dir: str) -> None:\n    checkpoint_path = Path(root_dir) / \"checkpoint\" / f\"{MODEL_NAME}.pt\"\n    components_path = Path(root_dir) / \"common\" / \"components.cif\"\n    components_pkl_path = Path(root_dir) / \"common\" / \"components.cif.rdkit_mol.pkl\"\n    if not checkpoint_path.exists():\n        raise FileNotFoundError(\n            f\"Missing checkpoint: {checkpoint_path}.\\n\"\n            \"Please include checkpoint/protenix_base_default_v1.0.0.pt in PROTENIX_ROOT_DIR.\"\n        )\n    if not components_path.exists():\n        raise FileNotFoundError(\n            f\"Missing CCD file: {components_path}.\\n\"\n            \"Please include common/components.cif in PROTENIX_ROOT_DIR.\"\n        )\n    if not components_pkl_path.exists():\n        raise FileNotFoundError(\n            f\"Missing CCD cache: {components_pkl_path}.\\n\"\n            \"Please include common/components.cif.rdkit_mol.pkl in PROTENIX_ROOT_DIR.\"\n        )\n\n\ndef build_input_json(test_df: pd.DataFrame, json_path: str) -> None:\n    data = []\n    for _, row in test_df.iterrows():\n        seq = row[\"sequence\"]\n        target_id = row[\"target_id\"]\n        data.append(\n            {\n                \"name\": target_id,\n                \"covalent_bonds\": [],\n                \"sequences\": [\n                    {\n                        \"rnaSequence\": {\n                            \"sequence\": seq,\n                            \"count\": 1,\n                        }\n                    }\n                ],\n            }\n        )\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\n\ndef build_single_input_json(target_id: str, seq: str, json_path: str) -> None:\n    data = [\n        {\n            \"name\": target_id,\n            \"covalent_bonds\": [],\n            \"sequences\": [\n                {\n                    \"rnaSequence\": {\n                        \"sequence\": seq,\n                        \"count\": 1,\n                    }\n                }\n            ],\n        }\n    ]\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\n\ndef build_configs(input_json_path: str, dump_dir: str, model_name: str):\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from configs.configs_model_type import model_configs\n    from protenix.config.config import parse_configs\n\n    base_configs = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(target, patch):\n        for key, value in patch.items():\n            if isinstance(value, dict) and key in target and isinstance(target[key], dict):\n                deep_update(target[key], value)\n            else:\n                target[key] = value\n\n    deep_update(base_configs, model_configs[model_name])\n\n    arg_str = \" \".join(\n        [\n            f\"--model_name {model_name}\",\n            f\"--input_json_path {input_json_path}\",\n            f\"--dump_dir {dump_dir}\",\n            f\"--use_msa {USE_MSA}\",\n            f\"--use_template {USE_TEMPLATE}\",\n            f\"--use_rna_msa {USE_RNA_MSA}\",\n            f\"--sample_diffusion.N_sample {MODEL_N_SAMPLE}\",\n            f\"--seeds {SEED}\",\n        ]\n    )\n    configs = parse_configs(\n        configs=base_configs,\n        arg_str=arg_str,\n        fill_required_with_null=True,\n    )\n    return configs\n\n\ndef get_c1_mask(data: dict, atom_array) -> torch.Tensor:\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, \"centre_atom_mask\"):\n                centre_mask = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    centre_mask = centre_mask & (atom_array.is_rna)\n                return torch.from_numpy(centre_mask)\n            if hasattr(atom_array, \"atom_name\"):\n                if hasattr(atom_array, \"is_rna\"):\n                    return torch.from_numpy(\n                        (atom_array.atom_name == \"C1'\") & (atom_array.is_rna)\n                    )\n                return torch.from_numpy(atom_array.atom_name == \"C1'\")\n        except Exception:\n            pass\n\n    features = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in features:\n        return features[\"centre_atom_mask\"].long() == 1\n    return features[\"atom_to_tokatom_idx\"].long() == 12\n\n\ndef get_feature_c1_mask(data: dict) -> torch.Tensor:\n    features = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in features:\n        return features[\"centre_atom_mask\"].long() == 1\n    return features[\"atom_to_tokatom_idx\"].long() == 12\n\n\ndef coords_to_rows(target_id: str, seq: str, coords: np.ndarray) -> list[dict]:\n    rows = []\n    n_res = len(seq)\n    for i in range(n_res):\n        row = {\n            \"ID\": f\"{target_id}_{i + 1}\",\n            \"resname\": seq[i],\n            \"resid\": i + 1,\n        }\n        for s in range(N_SAMPLE):\n            if s < coords.shape[0] and i < coords.shape[1]:\n                x, y, z = coords[s, i]\n            else:\n                x, y, z = 0.0, 0.0, 0.0\n            row[f\"x_{s + 1}\"] = float(x)\n            row[f\"y_{s + 1}\"] = float(y)\n            row[f\"z_{s + 1}\"] = float(z)\n        rows.append(row)\n    return rows\n\n\ndef pad_samples(coords: np.ndarray, target_samples: int) -> np.ndarray:\n    if coords.shape[0] >= target_samples:\n        return coords\n    if coords.shape[0] == 0:\n        return np.zeros((target_samples, coords.shape[1], 3), dtype=coords.dtype)\n    repeat = target_samples - coords.shape[0]\n    extra = np.repeat(coords[:1], repeat, axis=0)\n    return np.concatenate([coords, extra], axis=0)\n\n\ndef split_sequence(seq: str, max_len: int, overlap: int) -> list[tuple[int, int, str]]:\n    if max_len <= 0:\n        raise ValueError(\"max_len must be > 0\")\n    if overlap >= max_len:\n        raise ValueError(\"overlap must be smaller than max_len\")\n    if len(seq) <= max_len:\n        return [(0, len(seq), seq)]\n    step = max_len - overlap\n    chunks = []\n    start = 0\n    while start < len(seq):\n        end = min(start + max_len, len(seq))\n        chunks.append((start, end, seq[start:end]))\n        if end == len(seq):\n            break\n        start += step\n    return chunks\n\n\ndef main() -> None:\n    test_csv, output_csv, code_dir, root_dir = resolve_paths()\n\n    if not os.path.isdir(code_dir):\n        raise FileNotFoundError(\n            f\"Missing PROTENIX_CODE_DIR: {code_dir}. Set PROTENIX_CODE_DIR to the repo path.\"\n        )\n\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    sys.path.append(code_dir)\n\n    ensure_required_files(root_dir)\n    seed_everything(SEED)\n\n    test_df = pd.read_csv(test_csv)\n    if TEST_TARGETS:\n        keep = {t.strip() for t in TEST_TARGETS.split(\",\") if t.strip()}\n        test_df = test_df[test_df[\"target_id\"].isin(keep)].reset_index(drop=True)\n        if test_df.empty:\n            raise ValueError(\"TEST_TARGETS did not match any target_id values\")\n\n    work_dir = Path(\"/kaggle/working\")\n    \n    work_dir.mkdir(parents=True, exist_ok=True)\n    input_json_path = str(work_dir / \"protenix_input.json\")\n    build_input_json(test_df, input_json_path)\n\n    chunks_dir = work_dir / \"chunks\"\n    chunks_dir.mkdir(parents=True, exist_ok=True)\n\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import InferenceRunner, update_gpu_compatible_configs, update_inference_configs\n\n    configs = build_configs(input_json_path, str(work_dir / \"outputs\"), MODEL_NAME)\n    configs = update_gpu_compatible_configs(configs)\n    runner = InferenceRunner(configs)\n\n    dataset = InferenceDataset(configs)\n    seq_by_id = dict(zip(test_df.target_id.tolist(), test_df.sequence.tolist()))\n\n    all_predictions = []\n\n    debug_path = Path(output_csv).with_suffix(\".debug.csv\")\n    debug_rows = []\n\n    for i in tqdm(range(len(dataset)), total=len(dataset)):\n        data, atom_array, error_message = dataset[i]\n        target_id = data.get(\"sample_name\", f\"sample_{i}\")\n        seq = seq_by_id.get(target_id, \"\")\n\n        if len(seq) > MAX_SEQ_LEN:\n            chunks = split_sequence(seq, MAX_SEQ_LEN, CHUNK_OVERLAP)\n            print(\n                f\"[{target_id}] long sequence: {len(seq)} -> {len(chunks)} chunks \"\n                f\"(max_len={MAX_SEQ_LEN}, overlap={CHUNK_OVERLAP})\"\n            )\n            full_coords = np.zeros((N_SAMPLE, len(seq), 3), dtype=np.float32)\n            for chunk_idx, (start, end, subseq) in enumerate(chunks):\n                chunk_id = f\"{target_id}__chunk{chunk_idx:03d}\"\n                chunk_json_path = str(chunks_dir / f\"{chunk_id}.json\")\n                build_single_input_json(chunk_id, subseq, chunk_json_path)\n                chunk_configs = build_configs(\n                    chunk_json_path, str(work_dir / \"outputs\"), MODEL_NAME\n                )\n                chunk_configs = update_gpu_compatible_configs(chunk_configs)\n                chunk_dataset = InferenceDataset(chunk_configs)\n                chunk_data, chunk_atom_array, chunk_error = chunk_dataset[0]\n\n                if chunk_error:\n                    print(f\"[{chunk_id}] data_error: {chunk_error.splitlines()[0]}\")\n                    debug_rows.append(\n                        {\n                            \"target_id\": target_id,\n                            \"seq_len\": len(seq),\n                            \"c1_count\": 0,\n                            \"is_rna_atoms\": -1,\n                            \"coord_abs_max\": 0.0,\n                            \"error\": chunk_error,\n                            \"split_mode\": \"chunked\",\n                            \"chunk_idx\": chunk_idx,\n                            \"chunk_start\": start,\n                            \"chunk_end\": end,\n                            \"chunk_len\": end - start,\n                        }\n                    )\n                    continue\n\n                new_configs = update_inference_configs(\n                    configs, chunk_data[\"N_token\"].item()\n                )\n                runner.update_model_configs(new_configs)\n\n                prediction = runner.predict(chunk_data)\n                coords = prediction[\"coordinate\"]\n                mask = get_c1_mask(chunk_data, chunk_atom_array)\n                c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n                if c1_count == 0:\n                    mask = get_feature_c1_mask(chunk_data)\n                    c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n                    print(\n                        f\"[{chunk_id}] C1' atoms: {c1_count} (seq_len={len(subseq)}) fallback\"\n                    )\n                else:\n                    print(\n                        f\"[{chunk_id}] C1' atoms: {c1_count} (seq_len={len(subseq)})\"\n                    )\n                is_rna_count = (\n                    int(chunk_atom_array.is_rna.sum())\n                    if hasattr(chunk_atom_array, \"is_rna\")\n                    else -1\n                )\n                coord_abs_max = float(coords.abs().max().item())\n                debug_rows.append(\n                    {\n                        \"target_id\": target_id,\n                        \"seq_len\": len(seq),\n                        \"c1_count\": c1_count,\n                        \"is_rna_atoms\": is_rna_count,\n                        \"coord_abs_max\": coord_abs_max,\n                        \"error\": \"\",\n                        \"split_mode\": \"chunked\",\n                        \"chunk_idx\": chunk_idx,\n                        \"chunk_start\": start,\n                        \"chunk_end\": end,\n                        \"chunk_len\": end - start,\n                    }\n                )\n                coords = coords[:, mask, :].detach().cpu().numpy()\n\n                if coords.shape[1] != len(subseq):\n                    padded = np.zeros((coords.shape[0], len(subseq), 3), dtype=np.float32)\n                    min_len = min(coords.shape[1], len(subseq))\n                    if min_len > 0:\n                        padded[:, :min_len, :] = coords[:, :min_len, :]\n                    coords = padded\n\n                coords = pad_samples(coords, N_SAMPLE)\n\n                full_coords[:, start:end, :] = coords[:, : end - start, :]\n                del prediction, coords, mask, chunk_data, chunk_atom_array, chunk_dataset\n                torch.cuda.empty_cache()\n                gc.collect()\n\n            rows = coords_to_rows(target_id, seq, full_coords)\n            all_predictions.extend(rows)\n            continue\n\n        if error_message:\n            print(f\"[{target_id}] data_error: {error_message.splitlines()[0]}\")\n            debug_rows.append(\n                {\n                    \"target_id\": target_id,\n                    \"seq_len\": len(seq),\n                    \"c1_count\": 0,\n                    \"is_rna_atoms\": -1,\n                    \"coord_abs_max\": 0.0,\n                    \"error\": error_message,\n                    \"split_mode\": \"full\",\n                    \"chunk_idx\": -1,\n                    \"chunk_start\": 0,\n                    \"chunk_end\": len(seq),\n                    \"chunk_len\": len(seq),\n                }\n            )\n            rows = coords_to_rows(\n                target_id,\n                seq,\n                np.zeros((0, 0, 3), dtype=np.float32),\n            )\n            all_predictions.extend(rows)\n            continue\n\n        new_configs = update_inference_configs(configs, data[\"N_token\"].item())\n        runner.update_model_configs(new_configs)\n\n        prediction = runner.predict(data)\n        coords = prediction[\"coordinate\"]\n        mask = get_c1_mask(data, atom_array)\n        c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n        if c1_count == 0:\n            mask = get_feature_c1_mask(data)\n            c1_count = int(mask.sum().item()) if hasattr(mask, \"sum\") else 0\n            print(\n                f\"[{target_id}] C1' atoms: {c1_count} (seq_len={len(seq)}) fallback\"\n            )\n        else:\n            print(f\"[{target_id}] C1' atoms: {c1_count} (seq_len={len(seq)})\")\n        is_rna_count = (\n            int(atom_array.is_rna.sum()) if hasattr(atom_array, \"is_rna\") else -1\n        )\n        coord_abs_max = float(coords.abs().max().item())\n        debug_rows.append(\n            {\n                \"target_id\": target_id,\n                \"seq_len\": len(seq),\n                \"c1_count\": c1_count,\n                \"is_rna_atoms\": is_rna_count,\n                \"coord_abs_max\": coord_abs_max,\n                \"error\": \"\",\n                \"split_mode\": \"full\",\n                \"chunk_idx\": -1,\n                \"chunk_start\": 0,\n                \"chunk_end\": len(seq),\n                \"chunk_len\": len(seq),\n            }\n        )\n        coords = coords[:, mask, :].detach().cpu().numpy()\n\n        if coords.shape[1] != len(seq):\n            padded = np.zeros((coords.shape[0], len(seq), 3), dtype=np.float32)\n            min_len = min(coords.shape[1], len(seq))\n            if min_len > 0:\n                padded[:, :min_len, :] = coords[:, :min_len, :]\n            coords = padded\n\n        coords = pad_samples(coords, N_SAMPLE)\n\n        rows = coords_to_rows(target_id, seq, coords)\n        all_predictions.extend(rows)\n        del prediction, coords, mask, data, atom_array\n        torch.cuda.empty_cache()\n        gc.collect()\n\n    if debug_rows:\n        pd.DataFrame(debug_rows).to_csv(debug_path, index=False)\n\n    sub = pd.DataFrame(all_predictions)\n    cols = [\"ID\", \"resname\", \"resid\"] + [\n        f\"{c}_{i}\" for i in range(1, N_SAMPLE + 1) for c in [\"x\", \"y\", \"z\"]\n    ]\n    coord_cols = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\n    sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n    sub[cols].to_csv(output_csv, index=False)\n\n    print(f\"Saved submission to {output_csv}\")\n\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}