{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"},{"sourceId":11427182,"sourceType":"datasetVersion","datasetId":7089520},{"sourceId":11427187,"sourceType":"datasetVersion","datasetId":7078578},{"sourceId":11427191,"sourceType":"datasetVersion","datasetId":7078576},{"sourceId":326780,"sourceType":"modelInstanceVersion","modelInstanceId":274359,"modelId":295259}],"dockerImageVersionId":30918,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys  # sys module\nimport os\nimport torch\nimport pandas as pd\nimport numpy as np\nimport csv\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T23:59:01.316107Z","iopub.execute_input":"2025-04-15T23:59:01.316466Z","iopub.status.idle":"2025-04-15T23:59:01.320799Z","shell.execute_reply.started":"2025-04-15T23:59:01.316444Z","shell.execute_reply":"2025-04-15T23:59:01.319955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load datasets:\ntrain_seq = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\ntrain_lab = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\nval_seq = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_sequences.csv\")\nval_lab = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\")\ntest_seq =pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T23:59:01.325235Z","iopub.execute_input":"2025-04-15T23:59:01.325502Z","iopub.status.idle":"2025-04-15T23:59:01.591179Z","shell.execute_reply.started":"2025-04-15T23:59:01.325469Z","shell.execute_reply":"2025-04-15T23:59:01.590470Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Adding the user-defined module path\nsys.path.append(\"/kaggle/input/modules/\")\n\nfrom modules.dataset import RNAInferenceDataset\nfrom modules.models.baseline import RNABaselineBiLSTM\nfrom modules.utils.tm_score import compute_tm_scores_for_multiple_structures\n\n# ✅ Config\nclass Config:\n    def __init__(self):\n        self.batch_size = 32\n        self.device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n        self.model_path = \"/kaggle/input/rna_model_baseline/pytorch/default/1/rna_model_baseline.pth\"\n        self.test_dir = \"/kaggle/input/preprocessed-test\"\n        self.val_labels_path = \"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\"\n        self.sample_submission_path = \"/kaggle/input/stanford-rna-3d-folding/sample_submission.csv\"\n        self.submission_path = \"./submission.csv\"\n        self.hidden_dim = 128\n        self.num_layers = 2\n        self.dropout = 0.1\n\n# ✅ Load the model\ndef load_model(config):\n    model = RNABaselineBiLSTM(\n        input_dim=5,\n        hidden_dim=config.hidden_dim,\n        num_layers=config.num_layers,\n        dropout=config.dropout\n    )\n    checkpoint = torch.load(config.model_path, map_location=config.device)\n    model.load_state_dict(checkpoint)\n    model.to(config.device)\n    model.eval()\n    return model\n\n# ✅ Inference (updated for overall progress only)\ndef run_inference(model, dataloader, device):\n    predictions = []\n    total = len(dataloader.dataset)\n    with torch.no_grad():\n        with tqdm(total=total, desc=\"🔍 Running Inference\", unit=\"seq\") as pbar:\n            for batch in dataloader:\n                input_ids = batch[0].to(device)\n                output = model(input_ids)  # [B, L, 5, 3]\n                predictions.extend(output.cpu().numpy())\n                pbar.update(input_ids.size(0))\n    return predictions\n\n# ✅ Save submission\ndef save_submission_from_test_seq(sample_submission_path, target_ids, predictions, output_path):\n    df = pd.read_csv(sample_submission_path)\n\n    id2row = {}\n    tid2idx = {tid: i for i, tid in enumerate(target_ids)}\n    bad_ids = []\n\n    for _, row in df.iterrows():\n        full_id = row[\"ID\"]\n        target_id, resid = full_id.split(\"_\")\n        resid = int(resid)\n\n        if target_id not in tid2idx:\n            raise ValueError(f\"❌ target_id {target_id} not found\")\n\n        i = tid2idx[target_id]\n        coords = predictions[i][resid - 1]  # [5, 3]\n\n        if np.isnan(coords).any() or np.isinf(coords).any() or (coords == -1e+18).any():\n            bad_ids.append(full_id)\n            continue\n\n        res_char = row[\"resname\"]\n        new_row = [full_id, res_char, resid] + coords.flatten().tolist()\n        id2row[full_id] = new_row\n\n    if bad_ids:\n        print(f\"⚠️ {len(bad_ids)} entries have invalid coordinates and were skipped. Examples: {bad_ids[:5]}\")\n\n    header = ['ID', 'resname', 'resid'] + [f\"{axis}_{f}\" for f in range(1, 6) for axis in ['x', 'y', 'z']]\n    with open(output_path, \"w\", newline=\"\", encoding=\"utf-8\") as f:\n        writer = csv.writer(f, quoting=csv.QUOTE_NONE, escapechar=' ')\n        writer.writerow(header)\n        for full_id in df[\"ID\"]:\n            if full_id not in id2row:\n                raise ValueError(f\"❌ Missing ID: {full_id}\")\n            writer.writerow(id2row[full_id])\n\n    print(f\"✅ submission.csv saved successfully: {output_path} ({len(id2row)} rows)\")\n\n# ✅ Main 수정\ndef main():\n    config = Config()\n    config.sample_submission_path = \"/kaggle/input/stanford-rna-3d-folding/sample_submission.csv\"\n    config.submission_path = \"./submission.csv\"\n\n    print(\"📦 Loading model...\")\n    model = load_model(config)\n\n    print(\"📂 Loading .pt test data...\")\n    X = torch.load(os.path.join(config.test_dir, \"input_seqs.pt\"), map_location=config.device)\n    target_ids = torch.load(os.path.join(config.test_dir, \"target_ids.pkl\"), map_location=config.device)\n\n    try:\n        mask = torch.load(os.path.join(config.test_dir, \"masks.pt\"), map_location=config.device)\n        if mask.numel() == 0 or mask.shape[0] == 0:\n            raise ValueError(\"Empty mask detected\")\n    except:\n        print(\"⚠️ Using dummy mask\")\n        mask = torch.ones(X.shape[:2], dtype=torch.float32)\n\n    print(\"📊 Inference Dataset shapes:\")\n    print(\"X:\", X.shape)\n    print(\"mask:\", mask.shape)\n\n    test_dataset = RNAInferenceDataset(X, mask)\n    test_loader = DataLoader(test_dataset, batch_size=config.batch_size, num_workers=0)\n\n    print(\"🚀 Running inference...\")\n    predictions = run_inference(model, test_loader, config.device)\n\n    print(\"💾 Saving submission (test set)...\")\n    save_submission_from_test_seq(\n        config.sample_submission_path,\n        target_ids,\n        predictions,\n        config.submission_path\n    )\n\n    print(\"🎯 Submission ready for Kaggle upload!\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T23:59:01.592307Z","iopub.execute_input":"2025-04-15T23:59:01.592517Z","iopub.status.idle":"2025-04-15T23:59:01.972677Z","shell.execute_reply.started":"2025-04-15T23:59:01.592497Z","shell.execute_reply":"2025-04-15T23:59:01.971949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from config import Config  # 이미 되어 있다면 생략\nconfig = Config()\n\nconfig.test_dir = \"/kaggle/input/preprocessed-test\"  # 또는 너가 사용하는 validation 경로\n\ntarget_ids = torch.load(os.path.join(config.test_dir, \"target_ids.pkl\"))\n\nprint(\"📏 total target_ids:\", len(target_ids))\nprint(\"🧬 unique:\", len(set(target_ids)))\nprint(\"🧾 sample:\", target_ids[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T23:59:01.974155Z","iopub.execute_input":"2025-04-15T23:59:01.974410Z","iopub.status.idle":"2025-04-15T23:59:01.981095Z","shell.execute_reply.started":"2025-04-15T23:59:01.974390Z","shell.execute_reply":"2025-04-15T23:59:01.980145Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\n# Specify paths for submission files\nsubmission_path = './submission.csv'  # Path to your submission file\nsample_submission_path = '/kaggle/input/stanford-rna-3d-folding/sample_submission.csv'  # Path to sample submission file\n\n# Read both submission and sample submission files\nsub = pd.read_csv(submission_path)\nsample = pd.read_csv(sample_submission_path)\n\n# 1. Check the number of rows\nassert len(sub) == len(sample), f\"❌ Row count mismatch: submission ({len(sub)}) vs sample ({len(sample)})\"\nprint(f\"✅ Row count match: {len(sub)}\")\n\n# 2. Check for missing IDs\nmissing_ids = set(sample[\"ID\"]) - set(sub[\"ID\"])\nif len(missing_ids) > 0:\n    print(f\"❌ Missing IDs: {missing_ids}\")\nelse:\n    print(\"✅ All IDs are present in the submission file.\")\n\n# 3. Check column order and names\nexpected_cols = ['ID', 'resname', 'resid'] + [f\"{a}_{i}\" for i in range(1, 6) for a in 'xyz']\nassert list(sub.columns) == expected_cols, \"❌ Column order/names mismatch\"\nprint(\"✅ Column order and names match\")\n\n# 4. Check for NaN values\nassert sub.isnull().sum().sum() == 0, \"❌ NaN values found\"\nprint(\"✅ No NaN values\")\n\n# 5. Verify the data type of 'resid'\nassert pd.api.types.is_integer_dtype(sub['resid']), \"❌ 'resid' type error\"\nprint(\"✅ 'resid' column type is correct\")\n\n# 6. Verify coordinate data types (example: x_1 column)\nassert pd.api.types.is_float_dtype(sub['x_1']), \"❌ Coordinate type error\"\nprint(\"✅ Coordinate type is correct\")\n\n# 7. Verify no missing IDs\nassert set(sample[\"ID\"]) - set(sub[\"ID\"]) == set(), \"❌ Missing IDs found\"\nprint(\"✅ No missing IDs\")\n\n# 8. Compare data types of all columns\nprint(\"\\nSubmission file data types:\")\nprint(sub.dtypes)\n\nprint(\"\\nSample submission file data types:\")\nprint(sample.dtypes)\n\n# Check for any mismatched data types\nmismatched_columns = []\nfor col in sub.columns:\n    if sub[col].dtype != sample[col].dtype:\n        mismatched_columns.append(col)\n\nif mismatched_columns:\n    print(f\"❌ Mismatched data types in columns: {mismatched_columns}\")\nelse:\n    print(\"✅ All column data types match\")\n\nprint(\"🎉 Submission file format check complete!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T23:59:01.981998Z","iopub.execute_input":"2025-04-15T23:59:01.982219Z","iopub.status.idle":"2025-04-15T23:59:02.026039Z","shell.execute_reply.started":"2025-04-15T23:59:01.982200Z","shell.execute_reply":"2025-04-15T23:59:02.025354Z"}},"outputs":[],"execution_count":null}]}