{"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":12276181,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11278691,"sourceType":"datasetVersion","datasetId":7051341},{"sourceId":11279607,"sourceType":"datasetVersion","datasetId":7051942},{"sourceId":11469248,"sourceType":"datasetVersion","datasetId":7187409},{"sourceId":11856938,"sourceType":"datasetVersion","datasetId":7450304},{"sourceId":11952983,"sourceType":"datasetVersion","datasetId":7514818},{"sourceId":311741,"sourceType":"modelInstanceVersion","modelInstanceId":264400,"modelId":285488},{"sourceId":412196,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":336510,"modelId":357503}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install \\\n  \"/kaggle/input/fix-wheel/pyg_260_fixed/torch_scatter-2.1.2+pt26cu124-cp310-cp310-linux_x86_64.whl\" \\\n  \"/kaggle/input/fix-wheel/pyg_260_fixed/torch_sparse-0.6.18+pt26cu124-cp310-cp310-linux_x86_64.whl\" \\\n  \"/kaggle/input/fix-wheel/pyg_260_fixed/torch_cluster-1.6.3+pt26cu124-cp310-cp310-linux_x86_64.whl\" \\\n  \"/kaggle/input/fix-wheel/pyg_260_fixed/torch_spline_conv-1.2.2+pt26cu124-cp310-cp310-linux_x86_64.whl\" \\\n  \"/kaggle/input/fix-wheel/pyg_260_fixed/torch_geometric-2.6.1-py3-none-any.whl\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T02:28:29.827474Z","iopub.execute_input":"2025-05-26T02:28:29.827814Z","iopub.status.idle":"2025-05-26T02:28:33.436664Z","shell.execute_reply.started":"2025-05-26T02:28:29.827786Z","shell.execute_reply":"2025-05-26T02:28:33.435511Z"}},"outputs":[],"execution_count":null},{"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\nimport os\nimport sys","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-26T03:35:13.376725Z","iopub.execute_input":"2025-05-26T03:35:13.376963Z","iopub.status.idle":"2025-05-26T03:35:17.615704Z","shell.execute_reply.started":"2025-05-26T03:35:13.376942Z","shell.execute_reply":"2025-05-26T03:35:17.614507Z"}},"outputs":[],"execution_count":null},{"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    \"epochs\": 10,\n    \"cos_epoch\": 5,\n    \"loss_power_scale\": 1.0,\n    \"model_config_path\": \"../working/configs/pairwise.yaml\",\n    \"max_cycles\": 1,\n    \"grad_clip\": 0.1,\n    \"gradient_accumulation_steps\": 1,\n    \"d_clamp\": 30,\n    \"max_len_filter\": 9999999,\n    \"structural_violation_epoch\": 50,\n    \"balance_weight\": False,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:29:31.049197Z","iopub.execute_input":"2025-05-18T09:29:31.049445Z","iopub.status.idle":"2025-05-18T09:29:31.065201Z","shell.execute_reply.started":"2025-05-18T09:29:31.049427Z","shell.execute_reply":"2025-05-18T09:29:31.064537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\")\ntest_data.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:29:31.065990Z","iopub.execute_input":"2025-05-18T09:29:31.066206Z","iopub.status.idle":"2025-05-18T09:29:31.107502Z","shell.execute_reply.started":"2025-05-18T09:29:31.066178Z","shell.execute_reply":"2025-05-18T09:29:31.106768Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\n\nclass RNADataset(Dataset):\n    def __init__(self,data):\n        self.data=data\n        self.tokens={nt:i for i,nt in enumerate('ACGU')}\n\n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        sequence=[self.tokens[nt] for nt in (self.data.loc[idx,'sequence'])]\n        sequence=np.array(sequence)\n        sequence=torch.tensor(sequence)\n\n\n\n\n        return {'sequence':sequence}\n\ntest_dataset=RNADataset(test_data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:29:31.108329Z","iopub.execute_input":"2025-05-18T09:29:31.108524Z","iopub.status.idle":"2025-05-18T09:29:31.113583Z","shell.execute_reply.started":"2025-05-18T09:29:31.108508Z","shell.execute_reply":"2025-05-18T09:29:31.112704Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import math # math 모듈 임포트 확인\nimport torch # torch 임포트 확인\nimport torch.nn as nn\nsys.path.append('/kaggle/input/ribonanzanet2/pytorch/alpha/1')\nfrom Network import RibonanzaNet, MultiHeadAttention # MultiHeadAttention 임포트 추가 (SimpleStructureModule에서 사용)\n# from torch.utils import checkpoint # checkpoint 임포트 (원본 forward 메소드에서 사용)\n# Script B 원본에 checkpoint 임포트가 명시적으로 없었으나,\n# 모델의 forward 메소드에서 사용하고 있으므로 필요합니다.\n# 만약 이 스크립트의 다른 부분에서 이미 import torch.utils.checkpoint as checkpoint 했다면 중복 불필요\ntry:\n    from torch.utils.checkpoint import checkpoint\nexcept ImportError:\n    # PyTorch < 1.11 에서는 이름이 다를 수 있으나, 최신 버전 기준\n    print(\"Warning: torch.utils.checkpoint.checkpoint import failed. Ensure PyTorch version is compatible or import manually.\")\n    def checkpoint(fn, *args, **kwargs): # Dummy checkpoint for environments where it might be missing\n        return fn(*args)\n\n\nclass SinusoidalPosEmb(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n\n    def forward(self, x):\n        device = x.device\n        half_dim = self.dim // 2\n        # emb = math.log(10000) / (half_dim - 1) # 원본 코드\n        # half_dim이 1일 경우 ZeroDivisionError 발생 가능. Script A처럼 수정\n        denominator = (half_dim - 1) if half_dim > 1 else 1.0\n        emb = math.log(10000) / denominator\n        emb = torch.exp(torch.arange(half_dim, device=device) * -emb)\n        emb = x[:, None] * emb[None, :]\n        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)\n        return emb\n\nclass finetuned_RibonanzaNet(RibonanzaNet):\n    def __init__(self, rnet_config, config, pretrained=False): # 여기서 config는 diffusion_config를 의미\n        rnet_config.dropout=0.1\n        rnet_config.use_grad_checkpoint=True # Script B 원본 설정 유지\n        super(finetuned_RibonanzaNet, self).__init__(rnet_config)\n        if pretrained:\n            # 실제 사용시 config.pretrained_weight_path가 유효한 경로인지 확인 필요\n            self.load_state_dict(torch.load(config.pretrained_weight_path,map_location='cpu'))\n\n        self.dropout=nn.Dropout(0.0)\n\n        decoder_dim=config.decoder_dim\n        # SimpleStructureModule 정의가 이 코드 블록 이후에 오므로, 여기서 사용 가능\n        self.structure_module=[SimpleStructureModule(d_model=decoder_dim, nhead=config.decoder_nhead,\n                 dim_feedforward=decoder_dim*4, pairwise_dimension=rnet_config.pairwise_dimension, dropout=0.0) for i in range(config.decoder_num_layers)]\n        self.structure_module=nn.ModuleList(self.structure_module)\n\n        self.xyz_embedder=nn.Linear(3,decoder_dim)\n        self.xyz_norm=nn.LayerNorm(decoder_dim)\n        self.xyz_predictor=nn.Linear(decoder_dim,3)\n\n        self.adaptor=nn.Sequential(nn.Linear(rnet_config.ninp,decoder_dim),nn.LayerNorm(decoder_dim))\n\n        self.distogram_predictor=nn.Sequential(nn.LayerNorm(rnet_config.pairwise_dimension),\n                                                nn.Linear(rnet_config.pairwise_dimension,40))\n\n        self.time_embedder=SinusoidalPosEmb(decoder_dim)\n\n        self.time_mlp=nn.Sequential(nn.Linear(decoder_dim,decoder_dim),\n                                    nn.ReLU(),\n                                    nn.Linear(decoder_dim,decoder_dim))\n        self.time_norm=nn.LayerNorm(decoder_dim)\n\n        self.distance2pairwise=nn.Linear(1,rnet_config.pairwise_dimension,bias=False)\n\n        self.pair_mlp=nn.Sequential(nn.Linear(rnet_config.pairwise_dimension,rnet_config.pairwise_dimension),\n                                    nn.ReLU(),\n                                    nn.Linear(rnet_config.pairwise_dimension,rnet_config.pairwise_dimension))\n\n        #hyperparameters for diffusion\n        self.n_times = config.n_times # diffusion_config에서 n_times 가져옴\n\n        beta_1, beta_T = config.beta_min, config.beta_max\n        betas = torch.linspace(start=beta_1, end=beta_T, steps=config.n_times)\n        \n        # Script A/B의 버퍼 정의를 통합 및 필요한 모든 버퍼 정의\n        self.register_buffer(\"betas\", betas, persistent=False)\n        self.register_buffer(\"sqrt_betas\", betas.sqrt(), persistent=False)\n        \n        alphas = 1.0 - betas\n        self.register_buffer(\"alphas\", alphas, persistent=False)\n        self.register_buffer(\"sqrt_alphas\", alphas.sqrt(), persistent=False)\n        \n        alpha_bars = torch.cumprod(alphas, dim=0)\n        self.register_buffer(\"alpha_bars\", alpha_bars, persistent=False) # Script A는 없지만, B는 sqrt_alpha_bars 사용\n        self.register_buffer(\"sqrt_alpha_bars\", alpha_bars.sqrt(), persistent=False)\n        \n        # Script A는 sqrt_1mabar, Script B는 sqrt_one_minus_alpha_bars. 이름 통일 또는 둘 다 정의\n        # sqrt_one_minus_alpha_bars는 sqrt(1-alpha_bars)와 동일\n        self.register_buffer(\"sqrt_one_minus_alpha_bars\", (1.0 - alpha_bars).sqrt(), persistent=False)\n\n        self.data_std=config.data_std # diffusion_config에서 data_std 가져옴\n\n    def custom(self, module): # checkpoint를 위한 래퍼\n        def custom_forward(*inputs):\n            # SimpleStructureModule은 입력을 하나로 받으므로, inputs[0]을 전달\n            if len(inputs) == 1 and isinstance(inputs[0], (list, tuple)):\n                 return module(inputs[0])\n            return module(*inputs)\n        return custom_forward\n\n    def embed_pair_distance(self,inputs):\n        pairwise_features,xyz=inputs\n        distance_matrix=xyz[:,None,:,:]-xyz[:,:,None,:]\n        # clip min 값을 Script A 처럼 1e-9 (매우 작은 값) 또는 Script B의 2로 설정\n        # Script B의 clip(2, 37**2)는 물리적 의미가 있을 수 있으므로 유지\n        distance_matrix=(distance_matrix**2).sum(-1).clamp(min=1e-9, max=37**2).sqrt() # Script A의 clamp(min=1e-9) 적용\n        distance_matrix=distance_matrix[:,:,:,None]\n        pairwise_features=pairwise_features+self.distance2pairwise(distance_matrix)\n        return pairwise_features\n\n    # Script B의 원래 forward 메소드는 학습/단일 스텝 예측용으로 보임. 추론에는 직접 사용 안 함.\n    # 필요하다면 유지하되, 여기서는 레이턴트 추출을 위한 수정에 집중.\n    # def forward(self,src,xyz,t): ... (원본 Script B의 forward) ...\n\n    def denoise(self,sequence_features,pairwise_features,xyz_current_noise,ts_current_step): # 인자 이름 Script A와 유사하게 변경\n        # xyz_current_noise: (N_SAMPLES, SeqLen, 3)\n        # ts_current_step: (N_SAMPLES,)\n        N = xyz_current_noise.shape[0] # N_SAMPLES\n\n        # sequence_features, pairwise_features는 (1, L, D) 형태일 수 있으므로 N에 맞게 확장\n        # (이미 get_embeddings에서 배치 1로 나왔다고 가정)\n        seq_f_expanded = sequence_features.expand(N, -1, -1)\n        pair_f_expanded = pairwise_features.expand(N, -1, -1, -1)\n\n        current_pair_f = self.embed_pair_distance([pair_f_expanded, xyz_current_noise])\n        adapted_seq_f = self.adaptor(seq_f_expanded) # (N, L, decoder_dim)\n        \n        time_encoding = self.time_embedder(ts_current_step).unsqueeze(1) # (N, 1, decoder_dim)\n\n        # Script A의 denoise 로직 참고\n        tgt = adapted_seq_f + self.xyz_embedder(xyz_current_noise) + time_encoding\n        tgt = self.xyz_norm(tgt)\n\n        tgt_after_time_mlp = tgt + self.time_mlp(tgt) # Script A에서는 tgt + self.time_mlp(time_encoding) 이었으나, 여기선 tgt 사용\n        tgt = self.time_norm(tgt_after_time_mlp) # Script A의 self.time_norm(tgt + self.time_mlp(time_encoding)) 대신 사용\n\n        for layer in self.structure_module:\n            # SimpleStructureModule의 forward는 튜플/리스트를 단일 인자로 받음\n            # (tgt, src_features, pairwise_features, xyz, mask) 순서\n            tgt = layer((tgt, adapted_seq_f, current_pair_f, xyz_current_noise, None))\n\n        final_tgt_features = tgt # 이것이 레이턴트 벡터 (N_SAMPLES, SeqLen, decoder_dim)\n        epsilon_pred = self.xyz_predictor(final_tgt_features) # (N_SAMPLES, SeqLen, 3)\n\n        return epsilon_pred, final_tgt_features\n\n\n    def extract(self, a, t, x_shape):\n        # a: (total_timesteps,) 예를 들어 self.alphas\n        # t: (N_SAMPLES,) 각 샘플의 현재 타임스텝 인덱스\n        # x_shape: (N_SAMPLES, SeqLen, 3) 노이즈/데이터의 형태\n        device = t.device # a가 GPU에 있을 수도 있고 CPU에 있을 수도 있으므로, t의 디바이스 사용\n        a_gathered = torch.gather(a.to(device), 0, t.long()) # t를 인덱스로 사용하기 위해 long 타입으로\n        # view 대신 reshape 사용, Script A 방식과 동일하게\n        return a_gathered.view(t.size(0), *((1,) * (len(x_shape) - 1)))\n\n    # scale_to_minus_one_to_one, reverse_scale_to_zero_to_one, make_noisy는 학습용. 추론에서는 직접 사용 안함.\n    # 필요시 유지.\n\n    def denoise_at_t(self, x_t, sequence_features, pairwise_features, timesteps_batch, t_val):\n        # x_t: (N_SAMPLES, SeqLen, 3)\n        # sequence_features: (1, SeqLen, D_seq) 또는 (N_SAMPLES, SeqLen, D_seq) - denoise 내부에서 expand됨\n        # pairwise_features: (1, SeqLen, SeqLen, D_pair) 또는 (N_SAMPLES, ...) - denoise 내부에서 expand됨\n        # timesteps_batch: (N_SAMPLES,) 현재 스텝 인덱스 (예: 999, 998, ...)\n        # t_val: 스칼라 값 (현재 스텝 인덱스)\n\n        # 노이즈 추가: 마지막 스텝(t_val=0)에서는 노이즈 0, 그 외에는 랜덤 노이즈 (Script A 방식)\n        noise_for_step = torch.randn_like(x_t) if t_val > 0 else torch.zeros_like(x_t)\n        \n        # 수정된 denoise 호출: epsilon_pred와 final_tgt_at_this_step 반환\n        epsilon_pred, final_tgt_at_this_step = self.denoise(sequence_features, pairwise_features, x_t, timesteps_batch)\n        \n        # 필요한 계수들 추출 (self.alphas, self.sqrt_alphas 등은 __init__에서 올바르게 register_buffer 되어야 함)\n        alpha_t = self.extract(self.alphas, timesteps_batch, x_t.shape)\n        sqrt_alpha_t = self.extract(self.sqrt_alphas, timesteps_batch, x_t.shape)\n        sqrt_1m_alpha_bar_t = self.extract(self.sqrt_one_minus_alpha_bars, timesteps_batch, x_t.shape) # 이름 일치\n        sqrt_beta_t = self.extract(self.sqrt_betas, timesteps_batch, x_t.shape)\n        \n        # Script A의 디노이징 공식 적용 (수치 안정성을 위해 1e-9 추가)\n        term_eps_coeff = (1.0 - alpha_t) / (sqrt_1m_alpha_bar_t + 1e-9)\n        x_denoised_contribution = x_t - term_eps_coeff * epsilon_pred\n        x_t_minus_1 = (1.0 / (sqrt_alpha_t + 1e-9)) * x_denoised_contribution + sqrt_beta_t * noise_for_step\n        \n        return x_t_minus_1, final_tgt_at_this_step\n\n\n    def sample(self, src, N): # N은 N_SAMPLES\n        # src: (1, SeqLen) 입력 시퀀스 토큰\n        L_actual = src.shape[1]\n        x_t = torch.randn((N, L_actual, 3), device=src.device) # 초기 노이즈\n        \n        # 초기 임베딩 (배치 크기 1로 생성됨)\n        sequence_features, pairwise_features = self.get_embeddings(src, torch.ones_like(src).long().to(src.device))\n        # sequence_features: (1, L, D_seq), pairwise_features: (1, L, L, D_pair)\n        \n        # distogram 계산 (Script B 원본 로직 유지, device 일관성 및 squeeze 방식 수정)\n        distogram_logits = self.distogram_predictor(pairwise_features) # (1, L, L, 40)\n        \n        # squeeze는 차원이 1인 경우에만 수행하도록 하여 N_SAMPLES > 1일 때 문제 방지\n        squeezed_distogram_logits = distogram_logits\n        if distogram_logits.shape[0] == 1 and distogram_logits.ndim > 3: # (1, L, L, 40) 같은 경우\n             squeezed_distogram_logits = distogram_logits.squeeze(0) # -> (L, L, 40)\n\n        distogram_values = squeezed_distogram_logits[:,:,2:40] * torch.arange(2, 40, device=src.device).float()\n        distogram = distogram_values.sum(-1) # (L,L)\n\n        final_step_tgt_features_for_global = None # 글로벌 피처를 위한 변수\n\n        for t_val in range(self.n_times - 1, -1, -1): # 높은 t에서 낮은 t로 진행 (DDPM 역방향)\n            timesteps_batch = torch.full((N,), t_val, device=src.device, dtype=torch.long)\n            \n            # denoise_at_t는 이제 두 값을 반환\n            # sequence_features, pairwise_features는 여기서 (1,L,D) 형태로 전달되어도\n            # denoise_at_t -> denoise 내부에서 N_SAMPLES에 맞게 expand됨.\n            x_t, current_tgt_features = self.denoise_at_t(x_t, sequence_features, pairwise_features, timesteps_batch, t_val)\n            \n            if t_val == 0: # 마지막 스텝의 피처를 저장\n                final_step_tgt_features_for_global = current_tgt_features # (N, L, DecoderDim)\n        \n        x_0 = x_t * self.data_std # 최종 좌표 스케일링\n\n        global_features = None\n        if final_step_tgt_features_for_global is not None:\n            # (N, SeqLen, DecoderDim) -> (N, DecoderDim) 시퀀스 길이에 대해 평균\n            global_features = torch.mean(final_step_tgt_features_for_global, dim=1)\n        else:\n            print(\"Warning: sample 메서드에서 final_step_tgt_features_for_global이 설정되지 않았습니다.\")\n\n        return x_0, distogram, global_features # 글로벌 피처 추가 반환\n\nclass SimpleStructureModule(nn.Module):\n    def __init__(self, d_model, nhead,\n                 dim_feedforward, pairwise_dimension, dropout=0.1, # Script B 원본 dropout=0.1\n                 ):\n        super(SimpleStructureModule, self).__init__()\n        # MultiHeadAttention 클래스가 from Network import * 로 가져와졌다고 가정\n        self.self_attn = MultiHeadAttention(d_model, nhead, d_model//nhead, d_model//nhead, dropout=dropout)\n\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout = nn.Dropout(dropout) # FFN 중간의 드롭아웃\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.dropout1 = nn.Dropout(dropout) # 어텐션 후 드롭아웃\n        self.dropout2 = nn.Dropout(dropout) # FFN 후 드롭아웃\n\n        self.pairwise2heads=nn.Linear(pairwise_dimension,nhead,bias=False)\n        self.pairwise_norm=nn.LayerNorm(pairwise_dimension)\n\n        self.activation = nn.GELU()\n\n    def custom(self, module): # 이 custom 메소드는 gradient checkpointing을 위해 사용된 것으로 보임\n        def custom_forward(*inputs): # 추론 시에는 checkpoint를 사용하지 않으므로 직접 호출과 동일\n            # SimpleStructureModule의 forward는 튜플/리스트 하나를 받으므로,\n            # inputs가 ( (tgt, src, ...), ) 형태일 수 있음.\n            if len(inputs) == 1 and isinstance(inputs[0], (list, tuple)):\n                return module(inputs[0])\n            # 또는 *inputs가 이미 풀어진 (tgt, src, ...) 형태일 수 있음.\n            # 이 경우엔 return module(*inputs) 가 맞지만, finetuned_RibonanzaNet의 denoise에서\n            # layer((...)) 형태로 호출하므로, module(inputs[0])이 더 적절해 보임.\n            # 호출 방식에 따라 조정 필요. 여기서는 layer( (internal_args_tuple) ) 로 가정.\n            return module(inputs[0]) # 혹은 module(*inputs) - 호출 방식 확인 필요\n\n        return custom_forward\n\n    def forward(self, input_tuple): # 인자 이름을 input_tuple로 명확히 함\n        # input_tuple: (tgt, adapted_seq_f, current_pair_f, xyz_current_noise, None)\n        tgt , src_features, pairwise_features, xyz, src_mask = input_tuple # pred_t 대신 xyz, src 대신 src_features 사용\n        \n        pairwise_bias=self.pairwise2heads(self.pairwise_norm(pairwise_features)).permute(0,3,1,2)\n        # src_mask는 기본적으로 None으로 전달됨. MultiHeadAttention에서 처리.\n\n        res=tgt\n        # self.self_attn의 네 번째 인자는 key_padding_mask (src_mask와 유사한 역할이지만 형태 다름) 또는 attn_mask\n        # Script A에서는 bias로 사용된 pairwise_bias를 attn_mask (mask 인자)로 전달.\n        # src_mask는 key_padding_mask 역할.\n        # MultiHeadAttention(query, key, value, mask=attn_mask, src_mask=key_padding_mask)\n        # Script A: self_attn(tgt, tgt, tgt, mask=bias, src_mask=mask)\n        # 여기서 bias는 pairwise_bias, mask는 key_padding_mask.\n        # 여기서는 src_mask가 key_padding_mask 역할.\n        tgt, attention_weights = self.self_attn(tgt, tgt, tgt, mask=pairwise_bias, src_mask=src_mask)\n        tgt = res + self.dropout1(tgt)\n        tgt = self.norm1(tgt)\n\n        res=tgt\n        tgt = self.linear2(self.dropout(self.activation(self.linear1(tgt))))\n        tgt = res + self.dropout2(tgt)\n        tgt = self.norm2(tgt)\n\n        return tgt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:29:31.115760Z","iopub.execute_input":"2025-05-18T09:29:31.116018Z","iopub.status.idle":"2025-05-18T09:29:31.184749Z","shell.execute_reply.started":"2025-05-18T09:29:31.115998Z","shell.execute_reply":"2025-05-18T09:29:31.184075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import yaml\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\ndiffusion_config=load_config_from_yaml(\"/kaggle/input/ribonanzanet2-ddpm-v2/diffusion_config.yaml\")\nrnet_config=load_config_from_yaml(\"/kaggle/input/ribonanzanet2/pytorch/alpha/1/pairwise.yaml\")\n\nmodel=finetuned_RibonanzaNet(rnet_config,diffusion_config).cuda()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"state_dict=torch.load(\"/kaggle/input/ribonanzanet2-ddpm-v2/RibonanzaNet-DDPM-v2.pt\",map_location='cpu')\n\n#get rid of module. from ddp state dict\nnew_state_dict={}\n\nfor key in state_dict:\n    new_state_dict[key[7:]]=state_dict[key]\n\nmodel.load_state_dict(new_state_dict)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:29:31.185640Z","iopub.execute_input":"2025-05-18T09:29:31.185841Z","iopub.status.idle":"2025-05-18T09:29:33.781221Z","shell.execute_reply.started":"2025-05-18T09:29:31.185824Z","shell.execute_reply":"2025-05-18T09:29:33.780266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np # 혹시 NumPy가 import 안 되어있을 경우를 위해 추가\n\nmodel.eval()\n# preds=[] # 기존 이름 대신 명확하게 변경\npreds_xyz_all_rnas = []       # 3D 좌표 (xyz) 저장용 리스트\n# preds_distograms_all_rnas = [] # distogram 저장용 리스트 (필요하다면)\npreds_global_features_all_rnas = [] # 글로벌 피처 저장용 리스트 (새로 추가)\n\nfor i in tqdm(range(len(test_dataset))):\n    src = test_dataset[i]['sequence'].long()\n    src = src.unsqueeze(0).cuda()\n    # target_id = test_data.loc[i,'target_id'] # 이 변수는 현재 루프 내에서 직접 사용되지는 않음\n\n    # tmp=[] # 사용되지 않음\n    # predicted_dm=[] # 사용되지 않음\n    # for _ in range(5): # 루프 불필요, sample 메소드가 N개의 샘플을 한 번에 생성\n\n    with torch.no_grad():\n        # model.sample의 반환값이 3개로 변경됨\n        # xyz_samples: (N_SAMPLES, SeqLen, 3)\n        # distogram_pred: (SeqLen, SeqLen) - sample 메소드 구현에 따라 다를 수 있음\n        # global_features_pred: (N_SAMPLES, FeatureDim)\n        xyz_samples, distogram_pred, global_features_pred = model.sample(src, 5) # N_SAMPLES=5로 고정\n\n    preds_xyz_all_rnas.append(xyz_samples.cpu().numpy())\n    # preds_distograms_all_rnas.append(distogram_pred.cpu().numpy()) # distogram도 저장하려면 주석 해제\n\n    if global_features_pred is not None:\n        preds_global_features_all_rnas.append(global_features_pred.cpu().numpy())\n    else:\n        # 글로벌 피처가 None으로 반환될 경우를 대비 (예: sample 메소드 내에서 생성 실패 시)\n        # 또는 각 RNA 샘플별로 None을 추가하거나, 빈 NumPy 배열을 추가할 수 있습니다.\n        # 여기서는 (N_SAMPLES, 0) 형태의 빈 배열을 추가하여 차원 수를 유지하도록 함 (실제로는 None이 더 적절할 수 있음)\n        # 또는 아래 저장 로직에서 None을 처리하도록 함\n        preds_global_features_all_rnas.append(None) # 또는 np.empty((5, 0)) 등 상황에 맞게","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:29:38.993218Z","iopub.execute_input":"2025-05-18T09:29:38.993651Z","iopub.status.idle":"2025-05-18T09:39:17.723803Z","shell.execute_reply.started":"2025-05-18T09:29:38.993615Z","shell.execute_reply":"2025-05-18T09:39:17.722805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport torch\nimport numpy as np\nfrom tqdm.auto import tqdm # tqdm 추가\n\n# --- 이전에 정의되었거나 로드되었다고 가정하는 변수들 ---\n# preds_xyz_all_rnas: DDPM 추론 결과 3D 좌표 리스트 (각 요소는 NumPy 배열 (5, SeqLen, 3))\n# preds_global_features_all_rnas: DDPM 추론 결과 글로벌 피처 리스트 (각 요소는 NumPy 배열 (5, FeatureDim) 또는 None)\n# test_dataset: RNADataset 인스턴스 (src_tokens 가져오기용)\n# test_data: Pandas DataFrame (원본 test_sequences.csv 로드한 것. original_rna_id, rna_sequence_str 가져오기용)\n# DEVICE: torch.device 설정 (예: torch.device('cuda' if torch.cuda.is_available() else 'cpu'))\n# -------------------------------------------------------------\n\n# ───────────────── 스케일러 파라미터 로드 함수 ─────────────────\ndef load_scaler_params(scaler_json_path):\n    \"\"\" 저장된 스케일러 파라미터(평균, 표준편차)를 로드하는 함수 \"\"\"\n    try:\n        with open(scaler_json_path, 'r') as f:\n            scaler_params = json.load(f)\n        means = {k: float(v) for k, v in scaler_params['means'].items()}\n        stds = {k: float(v) for k, v in scaler_params['stds'].items()}\n        print(f\"Scaler params loaded from {scaler_json_path}\")\n        return means, stds\n    except FileNotFoundError:\n        print(f\"Error: Scaler params file not found at {scaler_json_path}. Using default (mean 0, std 1).\")\n        return {'x': 0., 'y': 0., 'z': 0.}, {'x': 1., 'y': 1., 'z': 1.}\n    except Exception as e:\n        print(f\"Error loading scaler params: {e}. Using default (mean 0, std 1).\")\n        return {'x': 0., 'y': 0., 'z': 0.}, {'x': 1., 'y': 1., 'z': 1.}\n\n# ───────────────── 텐서 좌표 처리 함수 (센터링, 표준화) ─────────────────\ndef center_coordinates_tensor(coords_tensor):\n    \"\"\" 3D 좌표 텐서를 센터링합니다. NaN을 제외하고 평균을 계산합니다. \"\"\"\n    if coords_tensor.ndim == 2: # (SeqLen, 3)\n        coords_tensor_batched = coords_tensor.unsqueeze(0) # (1, SeqLen, 3)으로 만듦\n    else: # (N_SAMPLES, SeqLen, 3)\n        coords_tensor_batched = coords_tensor\n\n    centered_coords_list = []\n    for i in range(coords_tensor_batched.shape[0]): # 각 샘플에 대해 처리\n        sample_coords = coords_tensor_batched[i] # (SeqLen, 3)\n        valid_mask = ~torch.isnan(sample_coords).any(dim=1)\n        if valid_mask.sum() > 0:\n            mean_for_centering = sample_coords[valid_mask].mean(dim=0, keepdim=True) # (1, 3)\n        else:\n            mean_for_centering = torch.zeros((1, 3), dtype=sample_coords.dtype, device=sample_coords.device)\n        centered_coords_list.append(sample_coords - mean_for_centering)\n    \n    if not centered_coords_list: # 모든 샘플이 유효하지 않은 극단적 경우\n         return torch.full_like(coords_tensor_batched, float('nan'))\n\n    output_tensor = torch.stack(centered_coords_list)\n    if coords_tensor.ndim == 2 and output_tensor.shape[0] == 1 : # 원래 차원으로 복원\n        output_tensor = output_tensor.squeeze(0)\n    return output_tensor\n\ndef standardize_coordinates_tensor(coords_tensor, means_dict, stds_dict, epsilon=1e-8):\n    \"\"\" 3D 좌표 텐서를 Z-score 표준화합니다. \"\"\"\n    if coords_tensor.ndim == 2: # (SeqLen, 3)\n        single_sample = True\n        coords_tensor_batched = coords_tensor.unsqueeze(0) # (1, SeqLen, 3)으로 만듦\n    else: # (N_SAMPLES, SeqLen, 3)\n        single_sample = False\n        coords_tensor_batched = coords_tensor\n    \n    device = coords_tensor_batched.device\n    # means_dict와 stds_dict의 타입이 float임을 load_scaler_params에서 보장\n    means_tensor = torch.tensor([[means_dict['x'], means_dict['y'], means_dict['z']]], dtype=torch.float32, device=device) # (1, 1, 3)\n    stds_tensor = torch.tensor([[stds_dict['x'], stds_dict['y'], stds_dict['z']]], dtype=torch.float32, device=device) # (1, 1, 3)\n\n    standardized_tensor = (coords_tensor_batched - means_tensor) / (stds_tensor + epsilon)\n    \n    if single_sample:\n        standardized_tensor = standardized_tensor.squeeze(0) # (SeqLen, 3)\n    return standardized_tensor\n\n# ───────────────── 데이터 전처리 실행 ─────────────────\n# <<<< 중요: 실제 스케일러 파라미터 파일 경로로 수정하세요 >>>>\nSCALER_PARAMS_PATH = \"/kaggle/input/data-for-egnn/coordinate_scaler_params_gb_feature.v2.json\" # 예시 경로\nloaded_means, loaded_stds = load_scaler_params(SCALER_PARAMS_PATH)\n\n# 최종적으로 그래프로 변환할 데이터를 담을 리스트\nprocessed_data_for_graph_conversion = []\n\nprint(\"\\nDDPM 추론 결과에 대한 데이터 전처리 (센터링 및 표준화)를 시작합니다...\")\n# len(test_data)는 DDPM 추론 루프에서 사용한 RNA 개수와 동일해야 함\nfor i in tqdm(range(len(test_data)), desc=\"센터링 및 표준화 중\"):\n    # test_dataset으로부터 src_tokens 가져오기\n    # DDPM 추론 시 사용한 src와 동일한 것을 사용해야 함\n    src_tokens_for_entry = test_dataset[i]['sequence'].long() # (SeqLen,)\n\n    # test_data DataFrame으로부터 ID 및 시퀀스 문자열 가져오기\n    original_rna_id = test_data.loc[i, 'target_id']\n    rna_sequence_str = test_data.loc[i, 'sequence']\n\n    # DDPM 추론 결과 (NumPy 배열)를 PyTorch 텐서로 변환 및 DEVICE로 이동\n    xyz_samples_raw_np = preds_xyz_all_rnas[i] # (5, SeqLen, 3) NumPy\n    xyz_samples_raw = torch.from_numpy(xyz_samples_raw_np).float().to(DEVICE)\n\n    global_features_np = preds_global_features_all_rnas[i] # (5, FeatureDim) NumPy 또는 None\n    global_features_pred_tensor = None\n    if global_features_np is not None and isinstance(global_features_np, np.ndarray) and global_features_np.size > 0:\n        global_features_pred_tensor = torch.from_numpy(global_features_np).float().to(DEVICE)\n\n    # 1. 센터링 적용 (각 샘플에 대해)\n    xyz_samples_centered = center_coordinates_tensor(xyz_samples_raw) # (5, SeqLen, 3) Tensor on DEVICE\n\n    # 2. Z-score 표준화 적용 (각 샘플에 대해)\n    xyz_samples_standardized = standardize_coordinates_tensor(xyz_samples_centered, loaded_means, loaded_stds) # (5, SeqLen, 3) Tensor on DEVICE\n\n    # 각 샘플별로 처리된 데이터를 딕셔너리 형태로 저장\n    for sample_idx in range(xyz_samples_standardized.shape[0]): # 5번 반복\n        current_xyz_std_list = xyz_samples_standardized[sample_idx].cpu().tolist()\n\n        current_global_feature_list = None\n        if global_features_pred_tensor is not None and sample_idx < global_features_pred_tensor.shape[0]:\n            current_global_feature_list = global_features_pred_tensor[sample_idx].cpu().tolist()\n\n        processed_entry = {\n            \"id\": f\"{original_rna_id}_{sample_idx + 1}\",\n            \"sequence\": rna_sequence_str,\n            \"coords_processed\": current_xyz_std_list, # 센터링 및 표준화된 좌표\n            \"global_feature\": current_global_feature_list,\n            \"src_tokens\": src_tokens_for_entry.cpu().tolist() # 원본 토큰\n        }\n        processed_data_for_graph_conversion.append(processed_entry)\n\nprint(f\"총 {len(processed_data_for_graph_conversion)}개의 처리된 entry 생성 완료.\")\nprint(\"이제 `processed_data_for_graph_conversion` 리스트를 사용하여 그래프 변환 및 DataLoader 생성을 진행할 수 있습니다.\")\n\n# # (선택 사항) 중간 결과 확인 또는 저장 - 디버깅용\n# if processed_data_for_graph_conversion:\n#     print(\"\\n첫 번째 처리된 entry 예시:\")\n#     print(json.dumps(processed_data_for_graph_conversion[0], indent=2))\n#\n#     # output_json_filepath_debug = \"debug_predictions_processed_latent.json\"\n#     # with open(output_json_filepath_debug, 'w') as f:\n#     #     json.dump(processed_data_for_graph_conversion, f, indent=2)\n#     # print(f\"디버깅용 처리 데이터가 '{output_json_filepath_debug}'에 저장되었습니다.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:39:17.724883Z","iopub.execute_input":"2025-05-18T09:39:17.725220Z","iopub.status.idle":"2025-05-18T09:39:17.843688Z","shell.execute_reply.started":"2025-05-18T09:39:17.725188Z","shell.execute_reply":"2025-05-18T09:39:17.843043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#score val\nimport pandas as pd\nimport pandas.api.types\nimport os\nimport re\n\n# Function to parse TMscore output\ndef parse_tmscore_output(output):\n    result = {}\n\n    # Extract TM-score based on length of reference structure (second)\n    tm_score_match = re.findall(r\"TM-score=\\s+([\\d.]+)\", output)[1]\n    result['TM-score'] = float(tm_score_match) if tm_score_match else None\n\n    return result\n\ndef write_pdb_line(atom_name, atom_serial, residue_name, chain_id, residue_num, x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P'):\n    \"\"\"\n    Writes a single line of PDB format based on provided atom information. \n    \n    Args:\n        atom_name (str): Name of the atom (e.g., \"N\", \"CA\").\n        atom_serial (int): Atom serial number.\n        residue_name (str): Residue name (e.g., \"ALA\"). \n        chain_id (str): Chain identifier. \n        residue_num (int): Residue number. \n        x_coord (float): X coordinate.\n        y_coord (float): Y coordinate.\n        z_coord (float): Z coordinate.\n        occupancy (float, optional): Occupancy value (default: 1.0). \n        b_factor (float, optional): B-factor value (default: 0.0). \n    \n    Returns:\n        str: A single line of PDB string.\n    \"\"\"\n    line = f\"ATOM  {atom_serial:>5d}  {atom_name:<5s} {residue_name:<3s} {residue_num:>3d}    {x_coord:>8.3f}{y_coord:>8.3f}{z_coord:>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n\"\n    return line\n\ndef write2pdb(df, xyz_id, pdb_path):\n    resolved_cnt=0\n    with open(pdb_path, \"w\") as pdb_file:\n        for _, row in df.iterrows():\n            x_coord=row[f\"x_{xyz_id}\"]\n            y_coord=row[f\"y_{xyz_id}\"]\n            z_coord=row[f\"z_{xyz_id}\"]\n\n            if x_coord>-1e17 and y_coord>-1e17 and z_coord>-1e17:\n            #if True:\n                resolved_cnt+=1\n                pdb_line = write_pdb_line(\n                    atom_name=\"C1'\", \n                    atom_serial=int(row[\"resid\"]), \n                    residue_name=row['resname'], \n                    chain_id='0', \n                    residue_num=int(row[\"resid\"]), \n                    x_coord=x_coord, \n                    y_coord=y_coord, \n                    z_coord=z_coord,\n                    atom_type=\"C\"\n                )\n                pdb_file.write(pdb_line)\n    return resolved_cnt\n\n\ndef score(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str) -> float:\n    '''\n    Computes the TM-score between predicted and native RNA structures using USalign.\n\n    This function evaluates the structural similarity of RNA predictions to native structures\n    by computing the TM-score. It uses USalign, a structural alignment tool, to compare\n    the predicted structures with the native structures.\n\n    Workflow:\n    1. Copies the USalign binary to the working directory and grants execution permissions.\n    2. Extracts the `pdb_id` from the `ID` column of both the solution and submission DataFrames.\n    3. Iterates over each unique `pdb_id`, grouping the native and predicted structures.\n    4. Writes PDB files for native and predicted structures.\n    5. Runs USalign on each predicted-native pair and extracts the TM-score.\n    6. Computes the highest TM-score per target and returns aggregated results.\n\n    Args:\n        solution (pd.DataFrame): A DataFrame containing the native RNA structures.\n        submission (pd.DataFrame): A DataFrame containing the predicted RNA structures.\n        row_id_column_name (str): The name of the column containing unique row identifiers.\n\n    Returns:\n        tuple:\n            - results (list): The highest TM-score for each `pdb_id`.\n            - results_per_sub (list): TM-scores for each predicted-native pair.\n            - outputs (list): Raw output logs from USalign for debugging.\n    '''\n\n    os.system(\"cp /kaggle/input/usalign/USalign /kaggle/working/\")\n    os.system(\"sudo chmod u+x /kaggle/working//USalign\")\n\n\n    # Extract pdb_id from ID (pdb_resid)\n    solution[\"pdb_id\"] = solution[\"ID\"].apply(lambda x: x.split(\"_\")[0])\n    submission[\"pdb_id\"] = submission[\"ID\"].apply(lambda x: x.split(\"_\")[0])\n\n    #fix pdb_ids comment out later\n    # solution.loc[solution['pdb_id']==\"R1138v1\",'pdb_id']='R1138'\n    # solution.loc[solution['pdb_id']==\"R1117\",'pdb_id']='R1117v2'\n    \n    results=[]\n    outputs=[]\n    results_per_sub=[]\n    # Iterate through each pdb_id and generate PDB files for both clean and corrupted data\n    for pdb_id, group_native in solution.groupby(\"pdb_id\"):\n        group_predicted = submission[submission[\"pdb_id\"] == pdb_id]\n        #print(group_native,group_predicted)\n        # Define output file paths\n        # clean_pdb_path = os.path.join(output_folder, f\"{pdb_id}_C3_clean.pdb\")\n        # corrupted_pdb_path = os.path.join(output_folder, f\"{pdb_id}_C3_corrupted.pdb\")\n        native_pdb=f'native.pdb'\n        predicted_pdb=f'predicted.pdb'\n\n        all_scores=[]\n        for pred_cnt in range(1,6):\n            tmp=[]\n            for native_cnt in range(1,41):\n                # Write solution PDB\n                resolved_cnt=write2pdb(group_native, native_cnt, native_pdb)\n                \n                # Write predicted PDB\n                _=write2pdb(group_predicted, pred_cnt, predicted_pdb)\n\n                if resolved_cnt>0:\n                    command = f\"/kaggle/working/USalign {predicted_pdb} {native_pdb} -atom \\\" C1'\\\"\"\n                    output = os.popen(command).read()\n                    outputs.append(output)\n                    parsed_data = parse_tmscore_output(output)\n                    tmp.append(parsed_data['TM-score'])\n                    \n            all_scores.append(max(tmp))\n        # print(output)\n        # stop\n        print(pdb_id)\n        print(all_scores)\n        results_per_sub.append(all_scores)\n        results.append(max(all_scores))\n    \n    print(results)\n    #return sum(results)/len(results), outputs\n    return results, results_per_sub, outputs\n    #return outputs\n\nif 'R1107' in set(test_data['target_id']):\n    solution=pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\")\n    submission = pd.read_csv('/kaggle/input/stanford-rna-3d-folding/sample_submission.csv')  # <<< Update to your submission path\n\n    \n    scores,results_per_sub,outputs=score(solution,submission,'ID')\n    print(np.mean(scores))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:39:17.844465Z","iopub.execute_input":"2025-05-18T09:39:17.844756Z","iopub.status.idle":"2025-05-18T09:40:44.506028Z","shell.execute_reply.started":"2025-05-18T09:39:17.844728Z","shell.execute_reply":"2025-05-18T09:40:44.505343Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- 이전 Script B의 추론 및 전처리 루프 이후 ---\n# processed_data_for_graph_conversion 리스트가 준비된 상태라고 가정\n# 예시:\n# processed_data_for_graph_conversion = [\n#  { \"id\": \"RNA_A_1\", \"sequence\": \"AUCG...\", \"coords_processed\": [[...],[...]],\n#    \"global_feature\": [...], \"src_tokens\": [...] },\n#  ...\n# ]\n# --------------------------------------------------\n\nimport torch\nimport torch.nn.functional as F\nfrom torch_geometric.data import Data, Batch as GeomBatch\nfrom torch_geometric.loader import DataLoader as PyGDataLoader\nfrom typing import List, Dict, Tuple\nimport random\nimport math\nfrom tqdm.auto import tqdm # tqdm 추가\n\n# ──────────────────── 0. 상수 (필요시 값 조정) ────────────────────\nNT2I       = {'A': 0, 'C': 1, 'G': 2, 'U': 3, 'N': 0} # 'N'을 'A'와 동일하게 처리\nKNN_K      = 10  # k-NN의 k 값\nEDGE_DIM   = 15  # 4(oh_i) + 4(oh_j) + 1(bb) + 3(delta) + 1(dist) + 1(ri) + 1(rj)\nEXP_CF_DIM = 768 # 예상되는 conditioning_feature(global_feature) 차원\nDEVICE     = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ──────────────────── 1. Edge builder (k-NN) ───────────────\ndef build_edges_knn(\n    pos_tensor: torch.Tensor, # k-NN 구성에 사용할 좌표 (표준화된 좌표 또는 센터링된 좌표)\n    seq_str: str,\n    resid_list: List[int],    # 1-based resid 리스트\n    k: int = KNN_K\n) -> Tuple[torch.Tensor, torch.Tensor]:\n    L = pos_tensor.size(0)\n    if L == 0:\n        return torch.empty((2, 0), dtype=torch.long, device=pos_tensor.device), \\\n               torch.empty((0, EDGE_DIM), dtype=torch.float32, device=pos_tensor.device)\n\n    dmat = torch.cdist(pos_tensor, pos_tensor) # 거리 행렬\n    pairs = set()\n\n    # k-최근접 이웃 (자기 자신 제외)\n    for i_node_idx in range(L):\n        neighbor_indices = torch.arange(L, device=pos_tensor.device)\n        neighbor_indices = neighbor_indices[neighbor_indices != i_node_idx]\n        if len(neighbor_indices) == 0:\n            continue\n        distances_to_i = dmat[i_node_idx, neighbor_indices]\n        sorted_neighbor_indices_by_dist = neighbor_indices[torch.argsort(distances_to_i)]\n        effective_k = min(k, len(sorted_neighbor_indices_by_dist))\n        for j_local_idx in range(effective_k):\n            pairs.add((i_node_idx, sorted_neighbor_indices_by_dist[j_local_idx].item()))\n\n    # Backbone (i ↔ i+1, 양방향) 엣지 추가\n    for i_node_idx in range(L - 1):\n        pairs.add((i_node_idx, i_node_idx + 1))\n        pairs.add((i_node_idx + 1, i_node_idx))\n\n    send_nodes, receive_nodes, edge_features_list = [], [], []\n    length_minus_one_normalized = max(L - 1, 1) # 정규화를 위한 분모 (0으로 나누기 방지)\n\n    for i_node, j_node in pairs:\n        dev = pos_tensor.device\n        #send_nodes.append(i_node)\n        #receive_nodes.append(j_node)\n        base_i_char = seq_str[i_node]\n        base_j_char = seq_str[j_node]\n        # NT2I.get의 두 번째 인자는 키가 없을 때 반환할 기본값 (여기서는 'N'의 인덱스)\n        one_hot_i = F.one_hot(torch.tensor(NT2I.get(base_i_char.upper(), NT2I['N']), device=dev), num_classes=4).float()\n        one_hot_j = F.one_hot(torch.tensor(NT2I.get(base_j_char.upper(), NT2I['N']), device=dev), num_classes=4).float()\n        is_backbone_edge = torch.tensor([1.0 if abs(i_node - j_node) == 1 else 0.0], device=dev)\n        delta_coords = pos_tensor[i_node] - pos_tensor[j_node] # 방향 벡터\n        distance_val = torch.norm(delta_coords, p=2).unsqueeze(0) # L2 norm 거리\n        # resid_list는 1-based index\n        resid_i_normalized = torch.tensor([(resid_list[i_node] - 1) / length_minus_one_normalized], device=dev)\n        resid_j_normalized = torch.tensor([(resid_list[j_node] - 1) / length_minus_one_normalized], device=dev)\n        edge_features_list.append(torch.cat([one_hot_i, one_hot_j, is_backbone_edge,\n                                             delta_coords, distance_val, resid_i_normalized, resid_j_normalized]))\n\n    edge_index_tensor = (torch.tensor([send_nodes, receive_nodes], dtype=torch.long, device=pos_tensor.device)\n                         if pairs else torch.empty((2, 0), dtype=torch.long, device=pos_tensor.device))\n    edge_attr_tensor = (torch.stack(edge_features_list).to(device=pos_tensor.device, dtype=torch.float32) if edge_features_list\n                        else torch.empty((0, EDGE_DIM), dtype=torch.float32, device=pos_tensor.device))\n    return edge_index_tensor, edge_attr_tensor\n\n# ──────────────────── 2. processed_entry → PyG Data ─────────────\ndef processed_entry_to_pyg_data(entry: Dict) -> Data:\n    \"\"\" Script B의 processed_entry 딕셔너리를 PyTorch Geometric Data 객체로 변환합니다.\"\"\"\n    seq_str      = entry[\"sequence\"]\n    entry_id     = entry[\"id\"]\n    # coords_processed는 이미 NumPy 배열에서 tolist()된 리스트 상태\n    pos_std_list = entry[\"coords_processed\"] # 표준화된 좌표 (EGNN 입력용)\n    global_feature_list = entry[\"global_feature\"]\n    # src_tokens_list = entry[\"src_tokens\"] # 필요시 사용\n\n    L = len(seq_str)\n    if not (L > 0 and isinstance(pos_std_list, list) and len(pos_std_list) == L):\n         raise ValueError(f\"ID {entry_id}: 시퀀스(L={L})와 좌표(L={len(pos_std_list) if pos_std_list else 'None'}) 길이 불일치 또는 좌표 타입 오류.\")\n\n    # PyG Data 객체 생성 시 모든 텐서는 동일한 디바이스에 있어야 함 (DEVICE로 통일)\n    pos_tensor_for_graph = torch.tensor(pos_std_list, dtype=torch.float32, device=DEVICE)\n\n    resname_list = [s_char for s_char in seq_str] # 예: ['A', 'U', 'G', ...]\n    resid_list   = list(range(1, L + 1))      # 예: [1, 2, ..., L] (1-based)\n\n    # k-NN 그래프 엣지 및 엣지 특징 생성 (위에서 정의한 함수 사용)\n    # 여기서는 표준화된 좌표(pos_tensor_for_graph)로 k-NN을 만듭니다.\n    edge_i, edge_a = build_edges_knn(pos_tensor_for_graph, seq_str, resid_list)\n\n    # 글로벌 피처 (conditioning_feature) 처리\n    cf_tensor = torch.zeros(EXP_CF_DIM, device=DEVICE) # 기본값은 0으로 채워진 텐서\n    if global_feature_list is not None:\n        temp_cf_tensor = torch.tensor(global_feature_list, dtype=torch.float32, device=DEVICE)\n        temp_cf_tensor = temp_cf_tensor.view(-1) # 1D로 만듦\n        if temp_cf_tensor.shape[0] == EXP_CF_DIM:\n            cf_tensor = temp_cf_tensor\n        elif temp_cf_tensor.shape[0] < EXP_CF_DIM:\n            padding = torch.zeros(EXP_CF_DIM - temp_cf_tensor.shape[0], device=DEVICE)\n            cf_tensor = torch.cat([temp_cf_tensor, padding])\n        else: # EXP_CF_DIM보다 긴 경우 자르기\n            cf_tensor = temp_cf_tensor[:EXP_CF_DIM]\n    # cf_tensor는 (EXP_CF_DIM,) 형태. collate 시 (BatchSize, 1, EXP_CF_DIM)이 될 수 있으므로 Dataset에서 (1, Dim)으로 만듦\n\n    # 노드 특징 'x' 생성: One-hot(resname) + normalized resid\n    one_hot_node_features = torch.stack([\n        F.one_hot(torch.tensor(NT2I.get(ch.upper(), NT2I['N'])), num_classes=4)\n        for ch in resname_list\n    ]).float().to(DEVICE)\n    resid_norm_node_features = ((torch.tensor(resid_list, dtype=torch.float32, device=DEVICE) - 1)\n                             / max(L - 1, 1)).unsqueeze(1)\n    node_x_features = torch.cat([one_hot_node_features, resid_norm_node_features], dim=1) # (L, 5)\n\n    # PyG Data 객체 생성\n    # EGNN 추론만 하는 경우, 타겟 y는 필요 없을 수 있습니다.\n    data = Data(\n        x=node_x_features,                # (L, 5) 노드 특징\n        pos=pos_tensor_for_graph,         # (L, 3) 표준화된 좌표 (EGNN 입력)\n        edge_index=edge_i,                # (2, NumEdges)\n        edge_attr=edge_a,                 # (NumEdges, EDGE_DIM)\n        conditioning_feature=cf_tensor.unsqueeze(0), # (1, EXP_CF_DIM) 형태로 저장 (collate 용이)\n        id=entry_id,\n        sequence_str=seq_str,\n        num_nodes = L # num_nodes 명시적 추가\n        # resname, resid는 필요시 추가 저장 가능 (이미 x 생성에 사용됨)\n    )\n    return data\n\n# ──────────────────── 그래프 데이터 생성 실행 ────────────────────\nprint(f\"\\n처리된 entry 리스트를 PyG Data 객체로 변환합니다 (총 {len(processed_data_for_graph_conversion)}개)...\")\npyg_data_list_for_inference = []\nskipped_conversion_count = 0\nfor entry_dict_item in tqdm(processed_data_for_graph_conversion, desc=\"PyG Data 객체 변환 중\"):\n    try:\n        pyg_data_item = processed_entry_to_pyg_data(entry_dict_item)\n        pyg_data_list_for_inference.append(pyg_data_item)\n    except ValueError as ve:\n        # print(f\"Warning: ID {entry_dict_item.get('id', '?')} 변환 오류 (ValueError): {ve}. 건너뜁니다.\")\n        skipped_conversion_count += 1\n    except Exception as e_conv:\n        # print(f\"Warning: ID {entry_dict_item.get('id', '?')} 변환 중 알 수 없는 오류: {e_conv}. 건너뜁니다.\")\n        skipped_conversion_count += 1\n\nif skipped_conversion_count > 0:\n    print(f\"Warning: 총 {skipped_conversion_count}개의 entry가 PyG Data 객체 변환에 실패하여 건너뛰었습니다.\")\nprint(f\"총 {len(pyg_data_list_for_inference)}개의 PyG Data 객체 생성 완료.\")\n\n# ──────────────────── 3. 단순화된 Dataset, Sampler, DataLoader ─────────────\n\n# 3.1) Dataset (메모리 내 그래프 리스트 사용)\nclass InMemoryGraphDataset(torch.utils.data.Dataset):\n    def __init__(self, pyg_data_list: List[Data]):\n        super().__init__()\n        self.graphs = pyg_data_list\n        # Data 객체 생성 시 conditioning_feature를 (1, EXP_CF_DIM)으로 이미 만듦\n\n    def __len__(self):\n        return len(self.graphs)\n\n    def __getitem__(self, idx):\n        return self.graphs[idx]\n\n# 3.2) BucketBatchSampler (선택 사항, 추론 시 단순 순차 처리도 가능)\n#      메모리가 매우 다양한 길이의 그래프로 인해 문제가 될 경우에만 유용.\n#      여기서는 단순화를 위해 사용하지 않음. 필요시 이전 코드에서 가져와 사용.\n\n# 3.3) Collate Function (conditioning_feature 처리)\ndef collate_graphs_for_egnn(data_list: List[Data]) -> GeomBatch:\n    batch = GeomBatch.from_data_list(data_list) # PyG가 대부분 자동 처리\n    # Data 객체 생성 시 conditioning_feature를 (1, EXP_CF_DIM)으로 만들었으므로,\n    # 배치 객체에서는 (NumGraphsInBatch, 1, EXP_CF_DIM)이 됨.\n    # EGNN 모델이 (NumGraphsInBatch, EXP_CF_DIM)을 기대한다면 squeeze(1) 필요.\n    if hasattr(batch, 'conditioning_feature') and batch.conditioning_feature is not None:\n        if batch.conditioning_feature.dim() == 3 and batch.conditioning_feature.shape[1] == 1:\n            batch.conditioning_feature = batch.conditioning_feature.squeeze(1)\n        # 차원 불일치 시 경고 또는 에러 처리 추가 가능\n    return batch\n\n# 3.4) 단순화된 DataLoader 생성 함수\ndef make_simple_inference_loader(\n    pyg_data_list: List[Data],\n    batch_size: int = 32, # GPU 메모리 및 추론 효율 고려하여 설정\n    num_workers: int = 0\n) -> PyGDataLoader:\n    if not pyg_data_list:\n        print(\"Warning: DataLoader 생성을 위한 PyG 데이터 리스트가 비어있습니다. None을 반환합니다.\")\n        return None\n\n    dataset = InMemoryGraphDataset(pyg_data_list)\n    \n    loader = PyGDataLoader(\n        dataset,\n        batch_size=batch_size,\n        shuffle=False, # 추론 시에는 셔플 불필요\n        num_workers=num_workers,\n        pin_memory=False,\n        #pin_memory=(DEVICE.type == 'cuda'),\n        collate_fn=collate_graphs_for_egnn\n    )\n    print(f\"단순 추론 DataLoader 생성 완료: {len(dataset)}개 그래프, {len(loader)}개 배치 (배치크기 {batch_size})\")\n    return loader\n\n# ──────────────────── DataLoader 생성 및 테스트 (추론용) ────────────────────\ninference_final_loader = None\nif pyg_data_list_for_inference:\n    inference_final_loader = make_simple_inference_loader(\n        pyg_data_list_for_inference,\n        batch_size=16 # 예시 배치 크기, 실제 사용 시 조정\n    )\n\n    if inference_final_loader and len(inference_final_loader) > 0:\n        print(\"\\n생성된 추론 DataLoader의 첫 번째 배치 정보:\")\n        try:\n            first_inference_batch = next(iter(inference_final_loader))\n            # Data 객체들이 이미 DEVICE로 옮겨졌으므로, 배치는 자동으로 해당 DEVICE에 생성됨.\n            # first_inference_batch = first_inference_batch.to(DEVICE) # 필요시 명시적 이동\n            print(first_inference_batch)\n            print(f\"  배치 내 그래프 수: {first_inference_batch.num_graphs}\")\n            if hasattr(first_inference_batch, 'conditioning_feature') and first_inference_batch.conditioning_feature is not None:\n                print(f\"  배치 conditioning_feature shape: {first_inference_batch.conditioning_feature.shape}\")\n            if hasattr(first_inference_batch, 'x') and first_inference_batch.x is not None:\n                print(f\"  배치 노드 특징(x) shape: {first_inference_batch.x.shape}\")\n            if hasattr(first_inference_batch, 'pos') and first_inference_batch.pos is not None:\n                print(f\"  배치 좌표(pos) shape: {first_inference_batch.pos.shape}\")\n        except Exception as e_loader_test:\n            print(f\"DataLoader 테스트 중 오류: {e_loader_test}\")\n            import traceback\n            traceback.print_exc()\n    else:\n        print(\"추론 DataLoader가 비어있거나 생성되지 않았습니다.\")\nelse:\n    print(\"PyG Data 리스트가 비어있어 DataLoader를 생성할 수 없습니다.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T09:40:44.506897Z","iopub.execute_input":"2025-05-18T09:40:44.507217Z","iopub.status.idle":"2025-05-18T09:40:44.525794Z","shell.execute_reply.started":"2025-05-18T09:40:44.507184Z","shell.execute_reply":"2025-05-18T09:40:44.524932Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#import torch\n#import torch.nn as nn\n#import torch.nn.functional as F\n#from torch_scatter import scatter_softmax, scatter_add\n#import math\n\n# ───────────────── Tiny Attention Block ─────────────────\nclass TinyAttn(nn.Module):\n    \"\"\"\n    A compact attention mechanism.\n    It computes scaled dot-product attention for given query, key, and value vectors,\n    followed by a projection and a small MLP.\n    \"\"\"\n    def __init__(self, dim: int, heads: int = 4, drop: float = 0.1):\n        super().__init__()\n        assert dim % heads == 0, f\"Dimension ({dim}) must be divisible by heads ({heads})\"\n        self.h = heads # Number of attention heads\n        self.d_head = dim // heads # Dimension of each head\n\n        # Linear layer to project input to Q, K, V\n        self.qkv = nn.Linear(dim, dim * 3, bias=False)\n        # Output projection layer\n        self.proj = nn.Linear(dim, dim)\n        # Layer normalization\n        self.ln = nn.LayerNorm(dim)\n        # Small feed-forward network (MLP)\n        self.mlp = nn.Sequential(\n            nn.Linear(dim, dim * 2),\n            nn.SiLU(), # Swish activation\n            nn.Dropout(drop),\n            nn.Linear(dim * 2, dim)\n        )\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, h: torch.Tensor, src_idx: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            h (torch.Tensor): Input tensor of shape (E, dim), where E is number of edges/elements.\n            src_idx (torch.Tensor): Source node indices for scatter_softmax, shape (E,).\n                                    Used to group edges by their source node for softmax.\n        Returns:\n            torch.Tensor: Output tensor of shape (E, dim).\n        \"\"\"\n        E, dim = h.shape # E = Number of elements (e.g., edges), dim = feature dimension\n\n        # Project to Q, K, V and split\n        q, k, v = self.qkv(h).chunk(3, dim=-1) # Each (E, dim)\n\n        # Reshape for multi-head attention: (E, num_heads, head_dim)\n        q = q.view(E, self.h, self.d_head)\n        k = k.view(E, self.h, self.d_head)\n        v = v.view(E, self.h, self.d_head)\n\n        # Calculate scaled dot-product attention scores (logits)\n        # (q * k) performs element-wise multiplication\n        # .sum(dim=-1) sums across the head_dim\n        logits = (q * k).sum(dim=-1) / math.sqrt(self.d_head) # Shape: (E, num_heads)\n\n        # Apply softmax grouped by source node index to get attention weights\n        attn_weights = scatter_softmax(logits, src_idx, dim=0) # Shape: (E, num_heads)\n\n        # Apply attention weights to values\n        # attn_weights.unsqueeze(-1) gives (E, num_heads, 1)\n        # v is (E, num_heads, head_dim)\n        # Result is (E, num_heads, head_dim)\n        weighted_values = attn_weights.unsqueeze(-1) * v\n\n        # Concatenate heads and reshape back to (E, dim)\n        ctx = weighted_values.contiguous().view(E, dim)\n\n        # Output projection, residual connection, and MLP block\n        out = h + self.drop(self.proj(ctx)) # Apply projection and add residual\n        out = out + self.drop(self.mlp(self.ln(out))) # Apply MLP block with residual\n        return out\n\n# ───────── DeepGate + Single‑Layer FiLM ─────────\nclass DeepGateCrossFiLM(nn.Module):\n    \"\"\"\n    DeepGate module with a *single* FiLM projection layer.\n    It processes edge attributes through an up-stack of TinyAttn blocks,\n    then modulates these features using a conditioning vector (global feature) via FiLM,\n    followed by a tap-attention mechanism and a down-stack of TinyAttn blocks.\n    The 'bert_dim' parameter is replaced by 'conditioning_feature_dim'.\n    \"\"\"\n\n    def __init__(\n        self,\n        edge_dim: int = 15, # Dimension of input edge features\n        up_dims: list[int] = [16, 32, 64, 128], # Dimensions for the up-stack attention blocks\n        down_dims: list[int] = [128, 64, 32, 16, 8], # Dimensions for the down-stack attention blocks\n        tap_dim: int = 128, # Dimension at the \"tap\" point (output of up-stack, input to FiLM and down-stack)\n        conditioning_feature_dim: int = None, # Dimension of the new global conditioning feature\n        bert_dim: int = None, # Old parameter name for backward compatibility (will be overridden by conditioning_feature_dim if both provided)\n        mlp_out: int = 256, # Output dimension of the MLP that processes the conditioning feature for FiLM\n        attn_heads: int = 4, # Number of attention heads in TinyAttn blocks\n        attn_drop: float = 0.1, # Dropout rate in TinyAttn blocks\n        eps: float = 1e-8, # Epsilon for numerical stability\n        **_ignored, # Allows for backward-compatibility with existing instantiation arguments\n    ):\n        super().__init__()\n        assert tap_dim == up_dims[-1], \"tap_dim must match the last dimension in up_dims\"\n        assert tap_dim == down_dims[0], \"tap_dim must match the first dimension in down_dims\"\n        for d_val in up_dims + down_dims: # Check divisibility for all attention block dimensions\n            assert d_val % attn_heads == 0, f\"Dimension {d_val} is not divisible by heads {attn_heads}\"\n\n        # Determine the final dimension for the conditioning feature\n        final_cond_dim = 768 # Default, e.g., if no dimension is specified\n        if conditioning_feature_dim is not None:\n            final_cond_dim = conditioning_feature_dim\n        elif bert_dim is not None: # Check for old parameter name if new one isn't provided\n            final_cond_dim = bert_dim\n            print(f\"Warning: Argument 'bert_dim' (value: {bert_dim}) is deprecated for DeepGateCrossFiLM. \"\n                  f\"Please use 'conditioning_feature_dim'. The dimension has been set to {final_cond_dim}.\")\n\n        self.conditioning_dim_val = final_cond_dim # Store the actual dimension used for conditioning\n        self.eps = eps\n\n        # --- Up‑stack: Processes initial edge features ---\n        self.in_proj = nn.Linear(edge_dim, up_dims[0]) # Initial projection of edge features\n        self.up_blocks = nn.ModuleList()\n        self.up_proj = nn.ModuleList() # Linear projections between up-stack blocks\n        for i, d_val in enumerate(up_dims):\n            self.up_blocks.append(TinyAttn(dim=d_val, heads=attn_heads, drop=attn_drop))\n            if i < len(up_dims) - 1: # Add projection if not the last block\n                self.up_proj.append(nn.Linear(d_val, up_dims[i + 1]))\n\n        # --- Single‑layer MLP‑FiLM: Modulates features using the conditioning vector ---\n        # MLP to process the conditioning vector\n        self.mlp_film = nn.Linear(self.conditioning_dim_val, mlp_out)\n        self.ln_mlp_film = nn.LayerNorm(mlp_out)\n        self.swish_film = nn.SiLU()\n\n        # Linear layer to generate gamma and beta for FiLM from the processed conditioning vector\n        self.to_gamma_beta = nn.Linear(mlp_out, tap_dim * 2) # tap_dim for gamma, tap_dim for beta\n\n        # LayerNorm for the features being conditioned by FiLM\n        self.ln_film_cond_target = nn.LayerNorm(tap_dim)\n\n        # --- Tap‑Attention: Computes attention weights based on FiLM-conditioned features ---\n        self.Wq_tap = nn.Linear(tap_dim, tap_dim) # Query projection for tap-attention\n        self.Wk_tap = nn.Linear(tap_dim, tap_dim) # Key projection for tap-attention\n\n        # --- Down‑stack: Further processes FiLM-conditioned and tap-attended features ---\n        self.down_blocks = nn.ModuleList()\n        self.down_proj = nn.ModuleList() # Linear projections between down-stack blocks\n        for i, d_val in enumerate(down_dims):\n            self.down_blocks.append(TinyAttn(dim=d_val, heads=attn_heads, drop=attn_drop))\n            if i < len(down_dims) - 1: # Add projection if not the last block\n                self.down_proj.append(nn.Linear(d_val, down_dims[i + 1]))\n\n        # Final projection to get scalar gating weights (phi)\n        self.to_phi = nn.Sequential(\n            nn.LayerNorm(down_dims[-1]),\n            nn.Linear(down_dims[-1], 1)\n        )\n\n    def forward(self, edge_attr: torch.Tensor, edge_index: torch.Tensor,\n                conditioning_vec: torch.Tensor, edge_batch: torch.Tensor):\n        \"\"\"\n        Args:\n            edge_attr (torch.Tensor): Edge features, shape (num_edges, edge_dim).\n            edge_index (torch.Tensor): Edge connectivity, shape (2, num_edges).\n            conditioning_vec (torch.Tensor): Global conditioning vector for the batch,\n                                             shape (num_graphs_in_batch, conditioning_dim_val).\n            edge_batch (torch.Tensor): Maps each edge to its graph index in the batch, shape (num_edges,).\n        Returns:\n            torch.Tensor: Gating weights phi multiplied by tap-attention alpha, shape (num_edges,).\n        \"\"\"\n        src_nodes = edge_index[0] # Source nodes for each edge\n\n        # --- Up‑stack ---\n        h_up = self.in_proj(edge_attr) # (E, up_dims[0])\n        for i, block in enumerate(self.up_blocks):\n            h_up = block(h_up, src_nodes)\n            if i < len(self.up_proj):\n                h_up = self.up_proj[i](h_up)\n        h_tapped = h_up  # Features at the tap point, shape (E, tap_dim)\n\n        # --- FiLM Conditioning ---\n        # Expand graph-level conditioning_vec to edge-level\n        # conditioning_vec is (num_graphs, cond_dim), edge_batch is (E,)\n        # b_expanded will be (E, cond_dim)\n        b_expanded = conditioning_vec[edge_batch]\n\n        # Process conditioning vector through MLP\n        processed_b = self.swish_film(self.ln_mlp_film(self.mlp_film(b_expanded))) # (E, mlp_out)\n\n        # Generate gamma and beta for FiLM\n        gamma, beta = self.to_gamma_beta(processed_b).chunk(2, dim=-1) # Each (E, tap_dim)\n\n        # Apply FiLM: h_film = gamma * h_tapped + beta\n        # Add residual connection after FiLM and LayerNorm\n        h_film_modulated = gamma * h_tapped + beta\n        h_conditioned_by_film = h_tapped + self.swish_film(self.ln_film_cond_target(h_film_modulated)) # (E, tap_dim)\n\n        # --- Tap‑Attention ---\n        # Use FiLM-conditioned features for Tap-Attention queries and keys\n        q_tap = self.Wq_tap(h_conditioned_by_film) # (E, tap_dim)\n        k_tap = self.Wk_tap(h_conditioned_by_film) # (E, tap_dim)\n\n        # Calculate tap-attention logits\n        alpha_logits = (q_tap * k_tap).sum(dim=-1) / math.sqrt(h_conditioned_by_film.size(-1)) # (E,)\n        # Apply scatter_softmax to get attention weights per source node\n        alpha_tap = scatter_softmax(alpha_logits, src_nodes, dim=0) # (E,)\n\n        # --- Down‑stack ---\n        # Input to down-stack is the FiLM-conditioned features\n        h_down = h_conditioned_by_film\n        for i, block in enumerate(self.down_blocks):\n            h_down = block(h_down, src_nodes)\n            if i < len(self.down_proj):\n                h_down = self.down_proj[i](h_down)\n\n        # Project to scalar gating weights phi\n        phi_gate_weights = self.to_phi(h_down).squeeze(-1) # (E,)\n\n        # Final output: element-wise product of tap-attention weights and gating weights\n        return alpha_tap * phi_gate_weights # (E,)\n\n# ───────────────── EGNNCore ─────────────────\nclass EGNNCore(nn.Module):\n    \"\"\"\n    Core Equivariant Graph Neural Network (EGNN) layer.\n    Updates node positions based on messages passed along edges, gated by DeepGateCrossFiLM.\n    \"\"\"\n    def __init__(self, gate_module: DeepGateCrossFiLM, eps: float = 1e-8):\n        super().__init__()\n        self.gate = gate_module # The gating mechanism (DeepGateCrossFiLM)\n        self.eps  = eps # Epsilon for numerical stability\n\n    def forward(self, pos: torch.Tensor, edge_index: torch.Tensor,\n                edge_attr: torch.Tensor, conditioning_vec: torch.Tensor,\n                edge_batch: torch.Tensor):\n        \"\"\"\n        Args:\n            pos (torch.Tensor): Node positions, shape (num_nodes, 3).\n            edge_index (torch.Tensor): Edge connectivity, shape (2, num_edges).\n            edge_attr (torch.Tensor): Edge features, shape (num_edges, edge_dim).\n            conditioning_vec (torch.Tensor): Global conditioning vector for the batch,\n                                             shape (num_graphs_in_batch, conditioning_dim_val).\n            edge_batch (torch.Tensor): Maps each edge to its graph index, shape (num_edges,).\n        Returns:\n            torch.Tensor: Updated node positions, shape (num_nodes, 3).\n        \"\"\"\n        src_nodes, dst_nodes = edge_index # Source and destination nodes for each edge\n\n        # Calculate gating weights using the DeepGate module\n        # These weights determine the influence of each edge message\n        gate_weights = self.gate(edge_attr, edge_index, conditioning_vec, edge_batch) # (num_edges,)\n\n        # Calculate difference vectors between source and destination node positions\n        pos_diff = pos[src_nodes] - pos[dst_nodes] # (num_edges, 3)\n\n        # Normalize difference vectors to get unit direction vectors\n        norm_pos_diff = pos_diff.norm(dim=-1, keepdim=True).clamp(min=self.eps) # (num_edges, 1)\n        unit_pos_diff = pos_diff / norm_pos_diff # (num_edges, 3)\n\n        # Calculate message updates for node positions\n        # Messages are scaled unit vectors, weighted by gate_weights\n        # gate_weights.unsqueeze(-1) makes it (num_edges, 1) for broadcasting\n        messages = gate_weights.unsqueeze(-1) * unit_pos_diff # (num_edges, 3)\n\n        # Aggregate messages for each node using scatter_add\n        # Sum messages for all edges pointing to the same source node (or could be destination, depending on convention)\n        # Here, messages are aggregated at the source_nodes.\n        # dim_size=pos.size(0) ensures the output tensor has a size for all nodes, even isolated ones.\n        delta_pos_updates = scatter_add(messages, src_nodes, dim=0, dim_size=pos.size(0)) # (num_nodes, 3)\n\n        # Update node positions by adding the aggregated messages\n        return pos + delta_pos_updates\n\n# ───────────────── EGNN ─────────────────\nclass EGNN(nn.Module):\n    \"\"\"\n    Equivariant Graph Neural Network (EGNN) model.\n    Applies EGNNCore for a specified number of steps to refine node positions.\n    Edge attributes are dynamically updated based on current node positions in each step.\n    \"\"\"\n    def __init__(self, n_steps: int = 1, **kwargs): # kwargs will catch conditioning_feature_dim or bert_dim\n        super().__init__()\n        # Instantiate the gating module, passing all kwargs (including conditioning_feature_dim)\n        gate_module = DeepGateCrossFiLM(**kwargs)\n        self.core    = EGNNCore(gate_module=gate_module, eps=gate_module.eps)\n        self.n_steps = n_steps # Number of EGNN update steps\n        self.eps     = gate_module.eps # Epsilon from the gate module\n\n    def forward(self, batch):\n        \"\"\"\n        Args:\n            batch: PyTorch Geometric Batch object containing graph data.\n                   Expected attributes: pos, edge_index, edge_attr, conditioning_feature, batch (for edge_batch).\n        Returns:\n            torch.Tensor: Final node positions after n_steps of updates.\n        \"\"\"\n        pos           = batch.pos # Initial node positions\n        edge_index    = batch.edge_index # Edge connectivity\n        initial_edge_attr = batch.edge_attr # Initial edge attributes\n\n        # Use batch.conditioning_feature instead of batch.bert\n        if not hasattr(batch, 'conditioning_feature'):\n            raise AttributeError(\"Batch object must have a 'conditioning_feature' attribute for EGNN.\")\n        conditioning_vec = batch.conditioning_feature # Global conditioning vector for the batch\n\n        # `batch.batch` maps each node to its graph index in the batch.\n        # `edge_batch` maps each edge to its graph index.\n        # This is derived from the source node of each edge.\n        edge_batch_map    = batch.batch[edge_index[0]]\n\n        if edge_index.numel() == 0: # Handle cases with no edges (e.g., single-node graphs)\n            return pos\n\n        # Structure of edge_attr (total 15 dim assumed by default in rna_graph_knn.py):\n        #   - oh_i (4), oh_j (4), backbone (1)  -> First 9 features (static_A)\n        #   - delta_coords (3)                  -> Features 9, 10, 11 (dynamically updated)\n        #   - distance_val (1)                  -> Feature 12 (dynamically updated)\n        #   - resid_i_norm (1), resid_j_norm (1) -> Last 2 features (static_B)\n        # Slicing indices for static parts of edge_attr:\n        static_part_A = initial_edge_attr[:, :9]  # One-hot encodings and backbone flag\n        static_part_B = initial_edge_attr[:, 13:] # Normalized residue indices\n\n        current_pos = pos # Node positions to be updated in each step\n\n        for step in range(self.n_steps):\n            if step == 0:\n                # Use initial edge attributes for the first step\n                current_edge_attr = initial_edge_attr\n            else:\n                # Dynamically update edge attributes based on current node positions\n                src_nodes, dst_nodes = edge_index\n                new_delta_coords  = current_pos[src_nodes] - current_pos[dst_nodes] # (E, 3)\n                new_distance_val   = new_delta_coords.norm(dim=-1, keepdim=True).clamp(min=self.eps) # (E, 1)\n                # Reconstruct edge_attr with new delta and distance\n                current_edge_attr = torch.cat([static_part_A, new_delta_coords, new_distance_val, static_part_B], dim=-1)\n\n            # Apply the EGNNCore update\n            current_pos = self.core(current_pos, edge_index, current_edge_attr,\n                                    conditioning_vec, edge_batch_map)\n        return current_pos\n\n# ───────────────── 가중치 초기화 ─────────────────\ndef init_weights(module: nn.Module):\n    \"\"\"\n    Initializes weights for nn.Linear and nn.LayerNorm modules.\n    Uses Kaiming uniform for Linear weights and zeros for biases.\n    Uses ones for LayerNorm weights and zeros for biases.\n    \"\"\"\n    if isinstance(module, nn.Linear):\n        # Check if weight exists and is not None before initialization\n        if hasattr(module, 'weight') and module.weight is not None:\n            nn.init.kaiming_uniform_(module.weight, a=math.sqrt(5)) # Kaiming uniform for weights\n        if hasattr(module, 'bias') and module.bias is not None:\n            nn.init.zeros_(module.bias) # Zeros for biases\n    elif isinstance(module, nn.LayerNorm):\n        if hasattr(module, 'weight') and module.weight is not None :\n            nn.init.ones_(module.weight) # Ones for LayerNorm weights\n        if hasattr(module, 'bias') and module.bias is not None:\n            nn.init.zeros_(module.bias) # Zeros for LayerNorm biases\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\negnn_model = EGNN(\n    n_steps                 = 3,\n    edge_dim                = 15,\n    up_dims                 = [16, 32, 64, 128],\n    down_dims               = [128, 64, 32, 16, 8],\n    tap_dim                 = 128,\n    conditioning_feature_dim = 768,  # 'bert_dim'을 'conditioning_feature_dim'으로 변경\n    mlp_out                 = 256,\n    attn_heads              = 4,\n    attn_drop               = 0.1,\n    eps                     = 1e-8,\n).to(device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def unstandardize_coordinates(coords_std_tensor, means_dict, stds_dict):\n    \"\"\"\n    Z-score 표준화된 좌표 텐서를 원래 스케일(센터링된 상태)로 되돌립니다.\n    Args:\n        coords_std_tensor: Tensor of shape (SeqLen,3) or (N,SeqLen,3), Z-score 표준화된 값\n        means_dict: {'x':…, 'y':…, 'z':…}\n        stds_dict:  {'x':…, 'y':…, 'z':…}\n    Returns:\n        Tensor in same shape as coords_std_tensor, but 원래 스케일(센터링된 상태)로 복원됨.\n    \"\"\"\n    single = coords_std_tensor.ndim == 2\n    if single:\n        coords = coords_std_tensor.unsqueeze(0)  # (1,SeqLen,3)\n    else:\n        coords = coords_std_tensor            # (N,SeqLen,3)\n\n    device = coords.device\n    means = torch.tensor([[means_dict['x'], means_dict['y'], means_dict['z']]],\n                         dtype=torch.float32, device=device)  # (1,1,3)\n    stds  = torch.tensor([[stds_dict['x'], stds_dict['y'], stds_dict['z']]],\n                         dtype=torch.float32, device=device)  # (1,1,3)\n\n    # 역변환: x_original = x_std * std + mean\n    restored = coords * stds + means\n\n    if single:\n        return restored.squeeze(0)  # (SeqLen,3)\n    return restored  # (N,SeqLen,3)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#import json # JSON 로드 (스케일러 파라미터)\n#import numpy as np\n#import pandas as pd\n#import torch\n# from tqdm.auto import tqdm # 이미 위에서 import 되었을 수 있음\n\n\n\n\n# EGNN 모델 인스턴스 생성 및 가중치 로드\nprint(\"\\nEGNN 모델을 생성하고 가중치를 로드합니다...\")\n\negnn_model.eval()\n# <<<< 중요: 실제 학습된 EGNN 가중치 파일 경로로 수정하세요 >>>>\negnn_model_weights_path = \"/kaggle/input/weight/pytorch/default/1/3layer_knn_1dRMAE9align_svd_MAE_best_model_epoch_12.pth\" # 예시 경로\nif os.path.exists(egnn_model_weights_path):\n    \n    try:\n        egnn_model.load_state_dict(torch.load(egnn_model_weights_path, map_location=DEVICE))\n        print(f\"EGNN 모델 가중치를 성공적으로 로드했습니다: {egnn_model_weights_path}\")\n    except Exception as e_load:\n         print(f\"EGNN 모델 가중치 로드 중 오류 발생 ({egnn_model_weights_path}): {e_load}. 초기화된 가중치를 사용합니다.\")\nelse:\n    print(f\"Warning: EGNN 모델 가중치 파일({egnn_model_weights_path})을 찾을 수 없습니다. 초기화된 가중치로 추론합니다.\")\negnn_model.eval()\n\n\n# ───────────────── 스케일러 파라미터 로드 (Z-score 역변환용) ─────────────────\n# <<<< 중요: 실제 스케일러 파라미터 파일 경로로 수정하세요 >>>>\nSCALER_PARAMS_PATH = \"/kaggle/input/data-for-egnn/coordinate_scaler_params_gb_feature.v2.json\" # 예시 경로\nloaded_means, loaded_stds = load_scaler_params(SCALER_PARAMS_PATH) # 이 함수는 이전 코드에 정의되어 있어야 함\n\n\n# ───────────────── EGNN 추론 및 결과 수집 (Z-score 역변환만 적용) ─────────────────\nprint(\"\\nEGNN 모델 추론을 시작합니다...\")\n# 결과를 RNA ID별, 샘플 번호별로 저장 (값: (SeqLen, 3) NumPy 배열 - 센터링된 스케일, Z-score 역변환됨)\nrna_id_to_final_coords = {}\n\nif inference_final_loader: # DataLoader가 성공적으로 생성되었다면\n    # <<<< 중요: egnn_model이 정의되고 가중치가 로드되었다고 가정합니다. >>>>\n    # if 'egnn_model' not in locals():\n    #     print(\"Error: EGNN 모델('egnn_model')이 정의되지 않았습니다. 임시 더미 모델을 사용합니다.\")\n    #     # 임시 더미 EGNN 모델 (실제 모델로 교체 필수)\n    #     class DummyEGNN(torch.nn.Module):\n    #         def __init__(self): super().__init__(); self.fc = torch.nn.Linear(3,3)\n    #         def forward(self, data): return data.pos + torch.randn_like(data.pos) * 0.01 # 입력 pos에 약간의 노이즈\n    #     egnn_model = DummyEGNN().to(DEVICE)\n    # egnn_model.eval()\n\n\n    with torch.no_grad():\n        for batch_graph_data in tqdm(inference_final_loader, desc=\"EGNN 모델 추론 중\"):\n            batch_graph_data = batch_graph_data.to(DEVICE)\n\n            # EGNN 모델 추론 (출력은 \"센터링된 상태에서 Z-score 표준화된\" 좌표라고 가정)\n            corrected_pos_batch_std_centered = egnn_model(batch_graph_data) # (TotalNodesInBatch, 3)\n\n            data_list_from_batch = batch_graph_data.to_data_list()\n            current_node_idx_in_batch = 0\n            for single_graph_data_from_batch in data_list_from_batch:\n                num_nodes = single_graph_data_from_batch.num_nodes\n\n                pred_coords_std_centered_single = corrected_pos_batch_std_centered[\n                    current_node_idx_in_batch : current_node_idx_in_batch + num_nodes\n                ]\n                current_node_idx_in_batch += num_nodes\n\n                graph_id_full = single_graph_data_from_batch.id # \"originalRNAid_sampleIdx\"\n                original_rna_id_key, sample_idx_str = graph_id_full.rsplit('_', 1)\n\n                # Z-score 역변환만 수행 (결과는 여전히 센터링된 상태의 원래 스케일)\n                pred_coords_centered_original_scale = unstandardize_coordinates( # 이 함수는 이전 코드에 정의됨\n                    pred_coords_std_centered_single, loaded_means, loaded_stds\n                )\n\n                if original_rna_id_key not in rna_id_to_final_coords:\n                    rna_id_to_final_coords[original_rna_id_key] = [None] * 5 # 5개 샘플 공간\n                try:\n                    sample_idx_one_based = int(sample_idx_str)\n                    if 1 <= sample_idx_one_based <= 5:\n                        rna_id_to_final_coords[original_rna_id_key][sample_idx_one_based - 1] = pred_coords_centered_original_scale.cpu().numpy()\n                    else:\n                        print(f\"Warning: ID {graph_id_full}의 sample_idx({sample_idx_one_based})가 유효 범위를 벗어남.\")\n                except ValueError:\n                     print(f\"Warning: ID {graph_id_full}의 sample_idx_str ('{sample_idx_str}') 변환 불가.\")\nelse:\n    print(\"생성된 DataLoader가 없어 EGNN 추론을 건너뜁니다.\")\n\n# ───────────────── 최종 Submission CSV 파일 생성 ─────────────────\nprint(\"\\n최종 Submission CSV 파일을 생성합니다...\")\nfinal_submission_rows = []\n\n# test_data는 원본 test_sequences.csv를 로드한 Pandas DataFrame이어야 합니다.\n# 이전에 Script B에서 test_data를 로드했으므로 해당 변수를 사용합니다.\nif 'test_data' not in locals() or not isinstance(test_data, pd.DataFrame) or test_data.empty:\n    print(\"Error: 'test_data' DataFrame이 비어있거나 정의되지 않았습니다. Submission CSV를 생성할 수 없습니다.\")\n    print(\"'/kaggle/input/stanford-rna-3d-folding/test_sequences.csv'에서 test_data를 로드하려고 시도합니다.\")\n    test_data_path_example = \"/kaggle/input/stanford-rna-3d-folding/test_sequences.csv\"\n    if os.path.exists(test_data_path_example):\n        test_data = pd.read_csv(test_data_path_example)\n        print(f\"'{test_data_path_example}'에서 test_data를 성공적으로 로드했습니다.\")\n    else:\n        print(f\"Error: '{test_data_path_example}'를 찾을 수 없어 test_data를 로드할 수 없습니다. CSV 생성을 중단합니다.\")\n        test_data = pd.DataFrame() # 빈 DataFrame으로 설정하여 아래 루프를 건너뛰도록 함\n\nif not test_data.empty:\n    for i in range(len(test_data)):\n        target_id_from_csv = test_data.loc[i, 'target_id']\n        sequence_from_csv = test_data.loc[i, 'sequence']\n        seq_len = len(sequence_from_csv)\n\n        predicted_coord_sets_for_this_rna = rna_id_to_final_coords.get(target_id_from_csv)\n\n        if predicted_coord_sets_for_this_rna is None:\n            print(f\"Warning: RNA ID {target_id_from_csv} EGNN 예측 없음. NaN 좌표 사용.\")\n            predicted_coord_sets_for_this_rna = [np.full((seq_len, 3), np.nan)] * 5\n\n        for j_residue_idx in range(seq_len):\n            row_for_csv = [\n                f\"{target_id_from_csv}_{j_residue_idx + 1}\",\n                sequence_from_csv[j_residue_idx],\n                j_residue_idx + 1\n            ]\n            for k_sample_idx in range(5): # 5개 예측 샘플에 대해\n                coords_for_this_sample = None\n                # predicted_coord_sets_for_this_rna의 길이가 5이고, 각 요소가 배열 또는 None일 수 있음\n                if k_sample_idx < len(predicted_coord_sets_for_this_rna) and \\\n                   predicted_coord_sets_for_this_rna[k_sample_idx] is not None:\n                    coords_for_this_sample = predicted_coord_sets_for_this_rna[k_sample_idx]\n\n                if coords_for_this_sample is not None and \\\n                   isinstance(coords_for_this_sample, np.ndarray) and \\\n                   coords_for_this_sample.shape == (seq_len, 3) and \\\n                   j_residue_idx < coords_for_this_sample.shape[0] and \\\n                   not np.isnan(coords_for_this_sample[j_residue_idx]).any(): # 유효한 좌표인지 확인\n                    row_for_csv.extend(coords_for_this_sample[j_residue_idx])\n                else:\n                    row_for_csv.extend([np.nan, np.nan, np.nan])\n            final_submission_rows.append(row_for_csv)\n\n    submission_columns = ['ID', 'resname', 'resid']\n    for i_sample_num in range(1, 6): # 1부터 5까지\n        submission_columns.extend([f\"x_{i_sample_num}\", f\"y_{i_sample_num}\", f\"z_{i_sample_num}\"])\n\n    submission_df_final = pd.DataFrame(final_submission_rows, columns=submission_columns)\n    submission_output_path = 'submission.csv' # 최종 파일 이름\n    submission_df_final.to_csv(submission_output_path, index=False)\n    print(f\"\\n✅ 최종 Submission 파일 생성 완료: {submission_output_path}\")\n    print(\"Submission 파일 샘플:\")\n    print(submission_df_final.head())\nelse:\n    print(\"test_data DataFrame이 비어있거나 로드되지 않아 Submission CSV를 생성할 수 없습니다.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T09:34:56.814613Z","iopub.execute_input":"2025-05-28T09:34:56.814823Z","iopub.status.idle":"2025-05-28T09:34:56.897594Z","shell.execute_reply.started":"2025-05-28T09:34:56.814802Z","shell.execute_reply":"2025-05-28T09:34:56.896473Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\ndef kabsch_align(P, Q):\n    \"\"\"\n    Kabsch 알고리즘으로 P (예측 좌표)를 Q (실제 좌표)에 정렬합니다.\n\n    Args:\n        P (np.ndarray): 예측 좌표, shape (N, 3)\n        Q (np.ndarray): 실제 좌표, shape (N, 3)\n\n    Returns:\n        np.ndarray: Q에 정렬된 P 좌표, shape (N, 3)\n    \"\"\"\n    assert P.shape == Q.shape, \"P와 Q는 동일한 shape이어야 합니다.\"\n    P_cent = P - P.mean(axis=0)\n    Q_cent = Q - Q.mean(axis=0)\n    C = np.dot(P_cent.T, Q_cent)\n    V, S, Wt = np.linalg.svd(C)\n    d = np.sign(np.linalg.det(np.dot(V, Wt)))\n    D = np.diag([1, 1, d])\n    U = np.dot(np.dot(V, D), Wt)\n    P_aligned = np.dot(P_cent, U) + Q.mean(axis=0)\n    return P_aligned\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_ground_truth_coords_with_backbone_only(\n    rna_id,\n    validation_csv_path=\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\"\n):\n    \"\"\"\n    특정 RNA ID의 실제 구조를 백본 선과 함께 시각화합니다.\n    \"\"\"\n    import pandas as pd\n    import matplotlib.pyplot as plt\n    from mpl_toolkits.mplot3d import Axes3D\n\n    df = pd.read_csv(validation_csv_path)\n    df_rna = df[df[\"ID\"].str.startswith(rna_id)]\n    if df_rna.empty:\n        print(f\"❌ RNA ID {rna_id}에 대한 ground truth 데이터가 없습니다.\")\n        return\n\n    df_rna = df_rna.sort_values(\"resid\")\n    coords = df_rna[['x_1', 'y_1', 'z_1']].dropna().values\n    if coords.shape[0] < 2:\n        print(\"⚠️ 좌표 수 부족으로 백본 시각화 불가.\")\n        return\n\n    fig = plt.figure(figsize=(8, 6))\n    ax = fig.add_subplot(111, projection='3d')\n    ax.plot(coords[:, 0], coords[:, 1], coords[:, 2], color='red', linewidth=1.5, label='Ground Truth Backbone')\n    ax.scatter(coords[:, 0], coords[:, 1], coords[:, 2], color='red', s=20, alpha=0.7)\n\n    ax.set_title(f\"Ground Truth RNA Structure with Backbone\\n{rna_id}\")\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n    ax.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predicted_coords_with_backbone_aligned(\n    rna_id,\n    sample_idx=0,\n    validation_csv_path=\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\"\n):\n    \"\"\"\n    특정 RNA ID의 예측 좌표를 Kabsch 정렬 후 백본 포함 시각화합니다.\n    \"\"\"\n    import pandas as pd\n    import matplotlib.pyplot as plt\n    from mpl_toolkits.mplot3d import Axes3D\n\n    # 실제 좌표 불러오기\n    df = pd.read_csv(validation_csv_path)\n    df_rna = df[df[\"ID\"].str.startswith(rna_id)]\n    if df_rna.empty:\n        print(f\"❌ RNA ID {rna_id}에 대한 validation 좌표가 없습니다.\")\n        return\n\n    df_rna = df_rna.sort_values(\"resid\")\n    true_coords = df_rna[['x_1', 'y_1', 'z_1']].dropna().values\n\n    # 예측 좌표 가져오기\n    pred_coords = rna_id_to_final_coords.get(rna_id, [None]*5)[sample_idx]\n    if pred_coords is None or pred_coords.shape != true_coords.shape:\n        print(f\"❌ 예측 좌표가 없거나 shape 불일치: {rna_id}\")\n        return\n\n    # Kabsch 정렬 적용\n    aligned_pred = kabsch_align(pred_coords, true_coords)\n\n    # 시각화\n    fig = plt.figure(figsize=(8, 6))\n    ax = fig.add_subplot(111, projection='3d')\n    ax.plot(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n            color='blue', linewidth=1.5, label=f'Predicted (Aligned)')\n    ax.scatter(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n               color='blue', s=20, alpha=0.7)\n\n    ax.set_title(f\"Predicted RNA Structure with Backbone\\n{rna_id} (Sample {sample_idx + 1})\")\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n    ax.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 실제 구조 시각화\nvisualize_ground_truth_coords_with_backbone_only(\"R1107\")\n\n# 예측 구조 (Kabsch 정렬 포함)\nvisualize_predicted_coords_with_backbone_aligned(\"R1107\", sample_idx=0)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_kabsch_aligned_prediction_vs_ground_truth(\n    rna_id,\n    sample_idx=0,\n    validation_csv_path=\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\"\n):\n    \"\"\"\n    Kabsch 정렬을 적용한 예측 구조 vs 실제 구조를 백본 선 포함해 시각화합니다.\n    \"\"\"\n    import pandas as pd\n    import matplotlib.pyplot as plt\n    from mpl_toolkits.mplot3d import Axes3D\n\n    df = pd.read_csv(validation_csv_path)\n    df_rna = df[df[\"ID\"].str.startswith(rna_id)]\n    if df_rna.empty:\n        print(f\"❌ RNA ID {rna_id}에 대한 validation 좌표가 없습니다.\")\n        return\n\n    df_rna = df_rna.sort_values(\"resid\")\n    true_coords = df_rna[['x_1', 'y_1', 'z_1']].dropna().values\n\n    pred_coords = rna_id_to_final_coords.get(rna_id, [None]*5)[sample_idx]\n    if pred_coords is None or pred_coords.shape != true_coords.shape:\n        print(f\"❌ 예측 좌표가 없거나 shape이 일치하지 않습니다: {rna_id}\")\n        return\n\n    # 👉 Kabsch 정렬\n    aligned_pred = kabsch_align(pred_coords, true_coords)\n\n    # 시각화\n    fig = plt.figure(figsize=(10, 8))\n    ax = fig.add_subplot(111, projection='3d')\n\n    ax.plot(true_coords[:, 0], true_coords[:, 1], true_coords[:, 2],\n            color='red', linewidth=1.5, label='Ground Truth Backbone')\n    ax.scatter(true_coords[:, 0], true_coords[:, 1], true_coords[:, 2],\n               color='red', s=20, alpha=0.7)\n\n    ax.plot(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n            color='blue', linewidth=1.5, label=f'Predicted (Kabsch-aligned)')\n    ax.scatter(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n               color='blue', s=20, alpha=0.7)\n\n    ax.set_title(f\"Kabsch-aligned RNA Structure Comparison\\n{rna_id} (Sample {sample_idx + 1})\")\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n    ax.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_kabsch_aligned_prediction_vs_ground_truth(\"R1107\", sample_idx=0)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predicted_coords_with_labels_and_arrows(\n    rna_id,\n    sample_idx=0,\n    validation_csv_path=\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\"\n):\n    import pandas as pd\n    import matplotlib.pyplot as plt\n    from mpl_toolkits.mplot3d import Axes3D\n    import numpy as np\n\n    # validation 실제 좌표 로딩 (정렬용)\n    df = pd.read_csv(validation_csv_path)\n    df_rna = df[df[\"ID\"].str.startswith(rna_id)].sort_values(\"resid\")\n    true_coords = df_rna[['x_1', 'y_1', 'z_1']].dropna().values\n\n    pred_coords = rna_id_to_final_coords.get(rna_id, [None]*5)[sample_idx]\n    if pred_coords is None or pred_coords.shape != true_coords.shape:\n        print(f\"❌ 좌표 불일치 또는 없음: {rna_id}\")\n        return\n\n    # Kabsch 정렬\n    aligned_pred = kabsch_align(pred_coords, true_coords)\n    N = aligned_pred.shape[0]\n\n    # 시각화\n    fig = plt.figure(figsize=(10, 8))\n    ax = fig.add_subplot(111, projection='3d')\n\n    # 백본 선 + 점\n    ax.plot(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n            color='blue', linewidth=1.2, label='Predicted Backbone')\n    ax.scatter(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n               color='blue', s=20, alpha=0.8)\n\n    # 번호 라벨 추가\n    for i in range(N):\n        x, y, z = aligned_pred[i]\n        ax.text(x, y, z, f\"{i+1}\", size=6, color='black', alpha=0.8)\n\n    # 화살표 방향 표시 (quiver 사용)\n    for i in range(N - 1):\n        start = aligned_pred[i]\n        direction = aligned_pred[i + 1] - aligned_pred[i]\n        ax.quiver(\n            start[0], start[1], start[2],\n            direction[0], direction[1], direction[2],\n            color='blue', linewidth=0.5, arrow_length_ratio=0.5, alpha=1.0\n        )\n\n    ax.set_title(f\"Predicted RNA Structure with Backbone and Direction\\n{rna_id} (Sample {sample_idx + 1})\")\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n    ax.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predicted_coords_with_labels_and_arrows(\"R1107\", sample_idx=0)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_ground_truth_coords_with_labels_and_arrows(\n    rna_id,\n    validation_csv_path=\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\"\n):\n    \"\"\"\n    실제 RNA 구조를 백본 선, residue 번호 라벨, 방향 화살표와 함께 시각화합니다.\n    \"\"\"\n    import pandas as pd\n    import matplotlib.pyplot as plt\n    import numpy as np\n    from mpl_toolkits.mplot3d import Axes3D\n\n    df = pd.read_csv(validation_csv_path)\n    df_rna = df[df[\"ID\"].str.startswith(rna_id)].sort_values(\"resid\")\n    true_coords = df_rna[['x_1', 'y_1', 'z_1']].dropna().values\n\n    if true_coords.shape[0] < 2:\n        print(f\"❌ RNA ID {rna_id}의 실제 좌표가 부족합니다.\")\n        return\n\n    N = true_coords.shape[0]\n\n    fig = plt.figure(figsize=(10, 8))\n    ax = fig.add_subplot(111, projection='3d')\n\n    ax.plot(true_coords[:, 0], true_coords[:, 1], true_coords[:, 2],\n            color='red', linewidth=1.5, label='Ground Truth Backbone')\n    ax.scatter(true_coords[:, 0], true_coords[:, 1], true_coords[:, 2],\n               color='red', s=20, alpha=0.8)\n\n    # residue 번호 라벨\n    for i in range(N):\n        x, y, z = true_coords[i]\n        ax.text(x, y, z, f\"{i+1}\", size=6, color='black', alpha=0.7)\n\n    # 방향 화살표\n    for i in range(N - 1):\n        start = true_coords[i]\n        direction = true_coords[i + 1] - true_coords[i]\n        ax.quiver(\n            start[0], start[1], start[2],\n            direction[0], direction[1], direction[2],\n            color='red', linewidth=0.5, arrow_length_ratio=0.5, alpha=1.0\n        )\n\n    ax.set_title(f\"Ground Truth RNA Structure with Backbone\\n{rna_id}\")\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n    ax.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_ground_truth_coords_with_labels_and_arrows(\"R1107\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_prediction_vs_ground_truth_with_backbone_and_arrows(\n    rna_id,\n    sample_idx=0,\n    validation_csv_path=\"/kaggle/input/stanford-rna-3d-folding/validation_labels.csv\"\n):\n    \"\"\"\n    특정 RNA ID에 대해 예측 및 실제 구조를 함께 시각화합니다.\n    백본 선, residue 번호, 방향 화살표 포함.\n    \"\"\"\n    import pandas as pd\n    import matplotlib.pyplot as plt\n    import numpy as np\n    from mpl_toolkits.mplot3d import Axes3D\n\n    df = pd.read_csv(validation_csv_path)\n    df_rna = df[df[\"ID\"].str.startswith(rna_id)].sort_values(\"resid\")\n    true_coords = df_rna[['x_1', 'y_1', 'z_1']].dropna().values\n\n    pred_coords = rna_id_to_final_coords.get(rna_id, [None]*5)[sample_idx]\n    if pred_coords is None or pred_coords.shape != true_coords.shape:\n        print(f\"❌ 좌표 불일치 또는 없음: {rna_id}\")\n        return\n\n    aligned_pred = kabsch_align(pred_coords, true_coords)\n    N = aligned_pred.shape[0]\n\n    fig = plt.figure(figsize=(10, 8))\n    ax = fig.add_subplot(111, projection='3d')\n\n    # 실제 구조\n    ax.plot(true_coords[:, 0], true_coords[:, 1], true_coords[:, 2],\n            color='red', linewidth=1.5, label='Ground Truth Backbone')\n    ax.scatter(true_coords[:, 0], true_coords[:, 1], true_coords[:, 2],\n               color='red', s=20, alpha=0.7)\n\n    # 예측 구조\n    ax.plot(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n            color='blue', linewidth=1.5, label=f'Predicted Backbone (Sample {sample_idx+1})')\n    ax.scatter(aligned_pred[:, 0], aligned_pred[:, 1], aligned_pred[:, 2],\n               color='blue', s=20, alpha=0.7)\n\n    # residue 번호 (예측 기준)\n    for i in range(N):\n        x, y, z = aligned_pred[i]\n        ax.text(x, y, z, f\"{i+1}\", size=6, color='black', alpha=0.6)\n\n    # 방향 화살표: 예측 (파란 화살표)\n    for i in range(N - 1):\n        start = aligned_pred[i]\n        direction = aligned_pred[i + 1] - aligned_pred[i]\n        ax.quiver(\n            start[0], start[1], start[2],\n            direction[0], direction[1], direction[2],\n            color='blue', linewidth=0.5, arrow_length_ratio=0.2, alpha=1.0\n        )\n\n    # 방향 화살표: 실제 (회색 화살표)\n    for i in range(N - 1):\n        start = true_coords[i]\n        direction = true_coords[i + 1] - true_coords[i]\n        ax.quiver(\n            start[0], start[1], start[2],\n            direction[0], direction[1], direction[2],\n            color='red', linewidth=0.5, arrow_length_ratio=0.5, alpha=1.0\n        )\n\n    ax.set_title(f\"Predicted vs Ground Truth RNA Structure\\n{rna_id} (Sample {sample_idx + 1})\")\n    ax.set_xlabel(\"X\")\n    ax.set_ylabel(\"Y\")\n    ax.set_zlabel(\"Z\")\n    ax.legend()\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_prediction_vs_ground_truth_with_backbone_and_arrows(\"R1107\", sample_idx=0)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}