{"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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":87793,"databundleVersionId":12276181},{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":8318191,"datasetId":4459124,"databundleVersionId":8449353},{"sourceType":"datasetVersion","sourceId":7639698,"datasetId":4299272,"databundleVersionId":7736182}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport random\nimport pickle\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:28:54.945657Z","iopub.execute_input":"2026-03-05T07:28:54.946014Z","iopub.status.idle":"2026-03-05T07:28:54.950689Z","shell.execute_reply.started":"2026-03-05T07:28:54.945987Z","shell.execute_reply":"2026-03-05T07:28:54.949799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#set seed for everything\ntorch.manual_seed(0)\nnp.random.seed(0)\nrandom.seed(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:28:54.955327Z","iopub.execute_input":"2026-03-05T07:28:54.955543Z","iopub.status.idle":"2026-03-05T07:28:54.965272Z","shell.execute_reply.started":"2026-03-05T07:28:54.955524Z","shell.execute_reply":"2026-03-05T07:28:54.964538Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"config = {\n    \"seed\": 0,\n    \"cutoff_date\": \"2020-01-01\",\n    \"test_cutoff_date\": \"2022-05-01\",\n    \"max_len\": 384,\n    \"batch_size\": 1,\n    \"learning_rate\": 1e-4,\n    \"weight_decay\": 0.0,\n    \"mixed_precision\": \"bf16\",\n    \"model_config_path\": \"../working/configs/pairwise.yaml\",  # Adjust path as needed\n    \"epochs\": 10,\n    \"cos_epoch\": 5,\n    \"loss_power_scale\": 1.0,\n    \"max_cycles\": 1,\n    \"grad_clip\": 0.1,\n    \"gradient_accumulation_steps\": 1,\n    \"d_clamp\": 30,\n    \"max_len_filter\": 9999999,\n    \"min_len_filter\": 10, \n    \"structural_violation_epoch\": 50,\n    \"balance_weight\": False,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:28:54.966562Z","iopub.execute_input":"2026-03-05T07:28:54.966848Z","iopub.status.idle":"2026-03-05T07:28:54.976709Z","shell.execute_reply.started":"2026-03-05T07:28:54.966815Z","shell.execute_reply":"2026-03-05T07:28:54.975782Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get data and do some data processing¶\n","metadata":{"execution":{"iopub.status.busy":"2025-02-27T00:35:07.639563Z","iopub.execute_input":"2025-02-27T00:35:07.63984Z","iopub.status.idle":"2025-02-27T00:35:07.643454Z","shell.execute_reply.started":"2025-02-27T00:35:07.639817Z","shell.execute_reply":"2025-02-27T00:35:07.64259Z"}}},{"cell_type":"code","source":"# Load data\n\ntrain_sequences=pd.read_csv(\"/kaggle/input/competitions/stanford-rna-3d-folding-2/train_sequences.csv\",low_memory=False)\ntrain_labels=pd.read_csv(\"/kaggle/input/competitions/stanford-rna-3d-folding-2/train_labels.csv\",low_memory=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:28:54.977936Z","iopub.execute_input":"2026-03-05T07:28:54.978184Z","iopub.status.idle":"2026-03-05T07:29:05.595753Z","shell.execute_reply.started":"2026-03-05T07:28:54.978158Z","shell.execute_reply":"2026-03-05T07:29:05.595072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0])\ntrain_labels[\"pdb_id\"] ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:29:05.596962Z","iopub.execute_input":"2026-03-05T07:29:05.597233Z","iopub.status.idle":"2026-03-05T07:29:08.00222Z","shell.execute_reply.started":"2026-03-05T07:29:05.597214Z","shell.execute_reply":"2026-03-05T07:29:08.001375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"float('Nan')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:29:08.003061Z","iopub.execute_input":"2026-03-05T07:29:08.003342Z","iopub.status.idle":"2026-03-05T07:29:08.00791Z","shell.execute_reply.started":"2026-03-05T07:29:08.00332Z","shell.execute_reply":"2026-03-05T07:29:08.007175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_xyz=[]\n\nfor pdb_id in tqdm(train_sequences['target_id']):\n    df = train_labels[train_labels[\"pdb_id\"]==pdb_id]\n    #break\n    xyz=df[['x_1','y_1','z_1']].to_numpy().astype('float32')\n    xyz[xyz<-1e17]=float('Nan');\n    all_xyz.append(xyz)\n\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T07:29:08.008748Z","iopub.execute_input":"2026-03-05T07:29:08.008988Z","iopub.status.idle":"2026-03-05T08:18:24.560979Z","shell.execute_reply.started":"2026-03-05T07:29:08.008957Z","shell.execute_reply":"2026-03-05T08:18:24.560109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# filter the data\n# Filter and process data\nfilter_nan = []\nmax_len = 0\nfor xyz in all_xyz:\n    if len(xyz) > max_len:\n        max_len = len(xyz)\n\n    #fill -1e18 masked sequences to nans\n    \n    #sugar_xyz = np.stack([nt_xyz['sugar_ring'] for nt_xyz in xyz], axis=0)\n    filter_nan.append((np.isnan(xyz).mean() <= 0.5) & \\\n                      (len(xyz)<config['max_len_filter']) & \\\n                      (len(xyz)>config['min_len_filter']))\n\nprint(f\"Longest sequence in train: {max_len}\")\n\nfilter_nan = np.array(filter_nan)\nnon_nan_indices = np.arange(len(filter_nan))[filter_nan]\n\ntrain_sequences = train_sequences.loc[non_nan_indices].reset_index(drop=True)\nall_xyz=[all_xyz[i] for i in non_nan_indices]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.561939Z","iopub.execute_input":"2026-03-05T08:18:24.562196Z","iopub.status.idle":"2026-03-05T08:18:24.678917Z","shell.execute_reply.started":"2026-03-05T08:18:24.562175Z","shell.execute_reply":"2026-03-05T08:18:24.678004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#pack data into a dictionary\n\ndata={\n      \"sequence\":train_sequences['sequence'].to_list(),\n      \"temporal_cutoff\": train_sequences['temporal_cutoff'].to_list(),\n      \"description\": train_sequences['description'].to_list(),\n      \"all_sequences\": train_sequences['all_sequences'].to_list(),\n      \"xyz\": all_xyz\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.679852Z","iopub.execute_input":"2026-03-05T08:18:24.680189Z","iopub.status.idle":"2026-03-05T08:18:24.684722Z","shell.execute_reply.started":"2026-03-05T08:18:24.680159Z","shell.execute_reply":"2026-03-05T08:18:24.683884Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Split train data into train/val/test¶\nWe will simply do a temporal split, because that's how testing is done in structural biology in general (in actual blind tests)","metadata":{}},{"cell_type":"code","source":"# Split data into train and test\nall_index = np.arange(len(data['sequence']))\ncutoff_date = pd.Timestamp(config['cutoff_date'])\ntest_cutoff_date = pd.Timestamp(config['test_cutoff_date'])\ntrain_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) <= cutoff_date]\ntest_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) > cutoff_date and pd.Timestamp(d) <= test_cutoff_date]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.687079Z","iopub.execute_input":"2026-03-05T08:18:24.68727Z","iopub.status.idle":"2026-03-05T08:18:24.719038Z","shell.execute_reply.started":"2026-03-05T08:18:24.687253Z","shell.execute_reply":"2026-03-05T08:18:24.718065Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(data['sequence'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.720694Z","iopub.execute_input":"2026-03-05T08:18:24.720933Z","iopub.status.idle":"2026-03-05T08:18:24.725835Z","shell.execute_reply.started":"2026-03-05T08:18:24.720905Z","shell.execute_reply":"2026-03-05T08:18:24.725005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Train size: {len(train_index)}\")\nprint(f\"Test size: {len(test_index)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.726857Z","iopub.execute_input":"2026-03-05T08:18:24.727165Z","iopub.status.idle":"2026-03-05T08:18:24.73816Z","shell.execute_reply.started":"2026-03-05T08:18:24.727136Z","shell.execute_reply":"2026-03-05T08:18:24.737357Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get pytorch dataset¶","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom ast import literal_eval\n\ndef get_ct(bp,s):\n    ct_matrix=np.zeros((len(s),len(s)))\n    for b in bp:\n        ct_matrix[b[0]-1,b[1]-1]=1\n    return ct_matrix\n\nclass RNA3D_Dataset(Dataset):\n    def __init__(self,indices,data):\n        self.indices=indices\n        self.data=data\n        self.tokens={nt:i for i,nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.indices)\n    \n    def __getitem__(self, idx):\n\n        idx=self.indices[idx]\n        sequence=[self.tokens[nt] for nt in (self.data['sequence'][idx])]\n        sequence=np.array(sequence)\n        sequence=torch.tensor(sequence)\n\n        #get C1' xyz\n        xyz=self.data['xyz'][idx]\n        xyz=torch.tensor(np.array(xyz))\n\n\n        if len(sequence)>config['max_len']:\n            crop_start=np.random.randint(len(sequence)-config['max_len'])\n            crop_end=crop_start+config['max_len']\n\n            sequence=sequence[crop_start:crop_end]\n            xyz=xyz[crop_start:crop_end]\n        \n\n        return {'sequence':sequence,\n                'xyz':xyz}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.738947Z","iopub.execute_input":"2026-03-05T08:18:24.739217Z","iopub.status.idle":"2026-03-05T08:18:24.874426Z","shell.execute_reply.started":"2026-03-05T08:18:24.739198Z","shell.execute_reply":"2026-03-05T08:18:24.873363Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset=RNA3D_Dataset(train_index,data)\nval_dataset=RNA3D_Dataset(test_index,data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.875195Z","iopub.execute_input":"2026-03-05T08:18:24.875412Z","iopub.status.idle":"2026-03-05T08:18:24.891073Z","shell.execute_reply.started":"2026-03-05T08:18:24.875394Z","shell.execute_reply":"2026-03-05T08:18:24.890372Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.891985Z","iopub.execute_input":"2026-03-05T08:18:24.892281Z","iopub.status.idle":"2026-03-05T08:18:24.904185Z","shell.execute_reply.started":"2026-03-05T08:18:24.892259Z","shell.execute_reply":"2026-03-05T08:18:24.903273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import plotly.graph_objects as go\nimport numpy as np\n\n\n\n# Example: Generate an Nx3 matrix\nxyz = train_dataset[len(train_dataset)-1]['xyz']  # Replace this with your actual Nx3 data\nN = len(xyz)\n\n\nfor _ in range(2): #plot twice because it doesnt show up on first try for some reason\n    # Extract columns\n    x, y, z = xyz[:, 0], xyz[:, 1], xyz[:, 2]\n    \n    # Create the 3D scatter plot\n    fig = go.Figure(data=[go.Scatter3d(\n        x=x, y=y, z=z,\n        mode='markers',\n        marker=dict(\n            size=5,\n            color=z,  # Coloring based on z-value\n            colorscale='Viridis',  # Choose a colorscale\n            opacity=0.8\n        )\n    )])\n    \n    # Customize layout\n    fig.update_layout(\n        scene=dict(\n            xaxis_title=\"X\",\n            yaxis_title=\"Y\",\n            zaxis_title=\"Z\"\n        ),\n        title=\"3D Scatter Plot\"\n    )\n\nfig.show()\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.904888Z","iopub.execute_input":"2026-03-05T08:18:24.905108Z","iopub.status.idle":"2026-03-05T08:18:24.928776Z","shell.execute_reply.started":"2026-03-05T08:18:24.905085Z","shell.execute_reply":"2026-03-05T08:18:24.928071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader=DataLoader(train_dataset,batch_size=1,shuffle=True)\nval_loader=DataLoader(val_dataset,batch_size=1,shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.929511Z","iopub.execute_input":"2026-03-05T08:18:24.9297Z","iopub.status.idle":"2026-03-05T08:18:24.933314Z","shell.execute_reply.started":"2026-03-05T08:18:24.929684Z","shell.execute_reply":"2026-03-05T08:18:24.932592Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Get RibonanzaNet¶\nWe will add a linear layer to predict xyz of C1' atoms","metadata":{}},{"cell_type":"code","source":"! pip install einops\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:24.934019Z","iopub.execute_input":"2026-03-05T08:18:24.934255Z","iopub.status.idle":"2026-03-05T08:18:28.340987Z","shell.execute_reply.started":"2026-03-05T08:18:24.934228Z","shell.execute_reply":"2026-03-05T08:18:28.339951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\n\nsys.path.append(\"/kaggle/input/ribonanzanet2d-final\")\n\n\nfrom Network import *\nimport yaml\n\n\n\nclass Config:\n    def __init__(self, **entries):\n        self.__dict__.update(entries)\n        self.entries=entries\n\n    def print(self):\n        print(self.entries)\n\ndef load_config_from_yaml(file_path):\n    with open(file_path, 'r') as file:\n        config = yaml.safe_load(file)\n    return Config(**config)\n\n\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, config, pretrained=False):\n        config.dropout=0.1\n        super(finetuned_RibonanzaNet, self).__init__(config)\n        if pretrained:\n            self.load_state_dict(torch.load(\"/kaggle/input/ribonanzanet-weights/RibonanzaNet.pt\",map_location='cpu'))\n        # self.ct_predictor=nn.Sequential(nn.Linear(64,256),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(256,64),\n        #                                 nn.ReLU(),\n        #                                 nn.Linear(64,1)) \n        self.dropout=nn.Dropout(0.0)\n        self.xyz_predictor=nn.Linear(256,3)\n\n\n    \n    def forward(self,src):\n        \n        #with torch.no_grad():\n        sequence_features, pairwise_features=self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n\n\n        xyz=self.xyz_predictor(sequence_features)\n\n        return xyz","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:28.342125Z","iopub.execute_input":"2026-03-05T08:18:28.342447Z","iopub.status.idle":"2026-03-05T08:18:30.273479Z","shell.execute_reply.started":"2026-03-05T08:18:28.342415Z","shell.execute_reply":"2026-03-05T08:18:30.272556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model=finetuned_RibonanzaNet(load_config_from_yaml(\"/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml\"),pretrained=True).cuda()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:30.274397Z","iopub.execute_input":"2026-03-05T08:18:30.274894Z","iopub.status.idle":"2026-03-05T08:18:32.300532Z","shell.execute_reply.started":"2026-03-05T08:18:30.274868Z","shell.execute_reply":"2026-03-05T08:18:32.299804Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training loop¶\nwe will use dRMSD loss on the predicted xyz. the loss function is invariant to translations, rotations, and reflections. because dRMSD is invariant to reflections, it cannot distinguish chiral structures, so there may be better loss functions","metadata":{}},{"cell_type":"code","source":"def calculate_distance_matrix(X,Y,epsilon=1e-4):\n    return (torch.square(X[:,None]-Y[None,:])+epsilon).sum(-1).sqrt()\n\n\ndef dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\n\n    mask=~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n    if d_clamp is not None:\n        rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).clip(0,d_clamp**2)\n    else:\n        rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n\n    return rmsd.sqrt().mean()/Z\n\ndef local_dRMSD(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=30):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\n\n    mask=(~torch.isnan(gt_dm))*(gt_dm<d_clamp)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n\n\n    rmsd=torch.square(pred_dm[mask]-gt_dm[mask])+epsilon\n    # rmsd=(torch.square(pred_dm[mask]-gt_dm[mask])+epsilon).sqrt()/Z\n    #rmsd=torch.abs(pred_dm[mask]-gt_dm[mask])/Z\n    return rmsd.sqrt().mean()/Z\n\ndef dRMAE(pred_x,\n          pred_y,\n          gt_x,\n          gt_y,\n          epsilon=1e-4,Z=10,d_clamp=None):\n    pred_dm=calculate_distance_matrix(pred_x,pred_y)\n    gt_dm=calculate_distance_matrix(gt_x,gt_y)\n\n\n\n    mask=~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()]=False\n\n    rmsd=torch.abs(pred_dm[mask]-gt_dm[mask])\n\n    return rmsd.mean()/Z\n\nimport torch\n\ndef align_svd_mae(input, target, Z=10):\n    \"\"\"\n    Aligns the input (Nx3) to target (Nx3) using SVD-based Procrustes alignment\n    and computes RMSD loss.\n    \n    Args:\n        input (torch.Tensor): Nx3 tensor representing the input points.\n        target (torch.Tensor): Nx3 tensor representing the target points.\n    \n    Returns:\n        aligned_input (torch.Tensor): Nx3 aligned input.\n        rmsd_loss (torch.Tensor): RMSD loss.\n    \"\"\"\n    assert input.shape == target.shape, \"Input and target must have the same shape\"\n\n    #mask \n    mask=~torch.isnan(target.sum(-1))\n\n    input=input[mask]\n    target=target[mask]\n    \n    # Compute centroids\n    centroid_input = input.mean(dim=0, keepdim=True)\n    centroid_target = target.mean(dim=0, keepdim=True)\n\n    # Center the points\n    input_centered = input - centroid_input.detach()\n    target_centered = target - centroid_target\n\n    # Compute covariance matrix\n    cov_matrix = input_centered.T @ target_centered\n\n    # SVD to find optimal rotation\n    U, S, Vt = torch.svd(cov_matrix)\n\n    # Compute rotation matrix\n    R = Vt @ U.T\n\n    # Ensure a proper rotation (det(R) = 1, no reflection)\n    if torch.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt @ U.T\n\n    # Rotate input\n    aligned_input = (input_centered @ R.T.detach()) + centroid_target.detach()\n\n    # # Compute RMSD loss\n    # rmsd_loss = torch.sqrt(((aligned_input - target) ** 2).mean())\n\n    # rmsd_loss = torch.sqrt(((aligned_input - target) ** 2).mean())\n    \n    # return aligned_input, rmsd_loss\n    return torch.abs(aligned_input-target).mean()/Z","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:18:32.301332Z","iopub.execute_input":"2026-03-05T08:18:32.301528Z","iopub.status.idle":"2026-03-05T08:18:32.311896Z","shell.execute_reply.started":"2026-03-05T08:18:32.301511Z","shell.execute_reply":"2026-03-05T08:18:32.311185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport math\nfrom tqdm import tqdm\nfrom torch.amp import GradScaler, autocast\n\n# 1. 基础设置与规范化\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nepochs = 18\ncos_epoch = 35\naccum_steps = 1 # 原来的 batch_size，实为梯度累加步数\n\n# 规范化无穷大值\nbest_val_loss = float('inf') \nbest_preds =[]\n\noptimizer = torch.optim.Adam(model.parameters(), weight_decay=0.0, lr=0.0001)\n\n# 2. 混合精度 Scaler\nscaler = GradScaler('cuda')\n\n# 计算 T_max 时使用 math.ceil 确保除法准确\ntotal_steps_per_epoch = math.ceil(len(train_loader) / accum_steps)\nschedule = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, \n    T_max=(epochs - cos_epoch) * total_steps_per_epoch\n)\n\nfor epoch in range(epochs):\n    model.train()\n    tbar = tqdm(train_loader, desc=f\"Epoch {epoch + 1} Train\")\n    total_loss = 0\n    valid_batches = 0 # 记录有效 batch 数量，防止 NaN batch 影响 Loss 平均计算\n    oom = 0\n    \n    optimizer.zero_grad() # 确保进入循环前梯度清零\n\n    for idx, batch in enumerate(tbar):\n        sequence = batch['sequence'].to(device)\n        gt_xyz = batch['xyz'].to(device).squeeze()\n\n        # 3. 正确开启混合精度训练 (AMP)\n        with autocast(device_type=device.type, dtype=torch.float16):\n            pred_xyz = model(sequence).squeeze()\n\n        # 2. 将输出转换回 float32，确保 SVD 能够正常计算且数值稳定\n        pred_xyz_f32 = pred_xyz.float()\n        gt_xyz_f32 = gt_xyz.float() # 确保 GT 也是 float32\n\n        # 3. 在 autocast 外部计算 Loss\n        loss = dRMAE(pred_xyz_f32, pred_xyz_f32, gt_xyz_f32, gt_xyz_f32) + \\\n               align_svd_mae(pred_xyz_f32, gt_xyz_f32)\n\n        # 4. 优化 NaN / Inf 处理逻辑\n        if torch.isnan(loss) or torch.isinf(loss):\n            continue # 遇到 NaN 直接跳过当前 batch，不要 break 终止整个 epoch\n\n        # 缩放 loss 用于梯度累加\n        scaled_loss = loss / accum_steps\n        scaler.scale(scaled_loss).backward()\n\n        total_loss += loss.item()\n        valid_batches += 1\n\n        # 5. 执行优化器步骤（包含正确的 Scaler 和梯度裁剪逻辑）\n        if (idx + 1) % accum_steps == 0 or (idx + 1) == len(train_loader):\n            # 必须先 unscale，然后才能进行梯度裁剪 (Clip Grad Norm)\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            \n            # Step 并更新 Scaler\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            \n            # 只有在指定 Epoch 后才开始调整学习率\n            if (epoch + 1) > cos_epoch:\n                schedule.step()\n\n        # 实时更新 Tqdm 后缀\n        current_avg_loss = total_loss / valid_batches if valid_batches > 0 else 0\n        tbar.set_postfix(Loss=f\"{current_avg_loss:.4f}\", OOMs=oom)\n\n    # ================= 验证环节 =================\n    model.eval()\n    tbar_val = tqdm(val_loader, desc=f\"Epoch {epoch + 1} Val\")\n    val_preds =[]\n    val_loss = 0\n    valid_val_batches = 0\n    \n    for idx, batch in enumerate(tbar_val):\n        sequence = batch['sequence'].to(device)\n        gt_xyz = batch['xyz'].to(device).squeeze()\n\n        with torch.no_grad():\n            # 推理使用混合精度加速\n            with autocast(device_type=device.type, dtype=torch.float16):\n                pred_xyz = model(sequence).squeeze()\n            \n            # 同样转换回 float32 计算验证集的 loss\n            pred_xyz_f32 = pred_xyz.float()\n            gt_xyz_f32 = gt_xyz.float()\n            loss = dRMAE(pred_xyz_f32, pred_xyz_f32, gt_xyz_f32, gt_xyz_f32)\n            \n        if torch.isnan(loss) or torch.isinf(loss):\n            continue\n            \n        val_loss += loss.item()\n        valid_val_batches += 1\n        \n        # 最好保存 float16 或者 numpy 的 float32 以节省内存\n        val_preds.append([gt_xyz_f32.cpu().numpy(), pred_xyz_f32.cpu().numpy()])\n        \n    avg_val_loss = val_loss / valid_val_batches if valid_val_batches > 0 else float('inf')\n    print(f\"Epoch {epoch + 1} Val loss: {avg_val_loss:.4f}\")\n    \n    if avg_val_loss < best_val_loss:\n        best_val_loss = avg_val_loss\n        best_preds = val_preds\n        torch.save(model.state_dict(), 'RibonanzaNet-3D.pt')\n        print(f\"--> Saved new best model with Val Loss: {best_val_loss:.4f}\")\n\n# 最终保存\ntorch.save(model.state_dict(), 'RibonanzaNet-3D-final.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-05T08:54:40.66801Z","iopub.execute_input":"2026-03-05T08:54:40.668401Z","execution_failed":"2026-03-05T08:56:44.469Z"}},"outputs":[],"execution_count":null}]}