{"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":11403143,"sourceType":"competition"},{"sourceId":11031805,"sourceType":"datasetVersion","datasetId":6870779}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install biopython","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:45:57.330014Z","iopub.execute_input":"2025-03-22T15:45:57.330288Z","iopub.status.idle":"2025-03-22T15:46:03.409634Z","shell.execute_reply.started":"2025-03-22T15:45:57.330257Z","shell.execute_reply":"2025-03-22T15:46:03.408522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport random\nimport pickle\nimport os\nfrom tqdm import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nimport yaml\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.amp import GradScaler\n\n# 재현성을 위한 시드 설정\ntorch.manual_seed(0)\nnp.random.seed(0)\nrandom.seed(0)\n\n# 설정\nconfig = {\n    \"seed\": 0,\n    \"cutoff_date\": \"2020-01-01\",\n    \"test_cutoff_date\": \"2022-05-01\",\n    \"max_len\": 384,\n    \"batch_size\": 1,  # 배치 사이즈 2로 증가\n    \"learning_rate\": 1e-4,\n    \"weight_decay\": 0.0,\n    \"mixed_precision\": \"bf16\",\n    \"model_config_path\": \"../working/configs/pairwise.yaml\",\n    \"epochs\": 200,  # 에폭 수 200으로 증가\n    \"cos_epoch\": 150,  # 코사인 스케줄러 시작점 조정\n    \"loss_power_scale\": 1.0,\n    \"max_cycles\": 1,\n    \"grad_clip\": 1.0,\n    \"gradient_accumulation_steps\": 1,\n    \"d_clamp\": 30,\n    \"max_len_filter\": 9999999,\n    \"min_len_filter\": 10, \n    \"structural_violation_epoch\": 50,\n    \"balance_weight\": False,\n    \"msa_max_sequences\": 32,  # MSA 최대 시퀀스 수\n    \"msa_feat_dim\": 128,      # MSA 특성 차원\n    \"num_self_attn_layers\": 4,  # Self-attention 레이어 수\n    \"num_cross_attn_layers\": 4,  # Cross-attention 레이어 수\n    \"num_structure_module_layers\": 8,  # 구조 모듈 레이어 수\n    \"dropout\": 0.25,  # 드롭아웃 증가 (0.1 → 0.25)\n    \"n_heads\": 8,\n    \"convert_to_rna\": False  # T를 U로 변환 않도록 변경\n}","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:46:03.410743Z","iopub.execute_input":"2025-03-22T15:46:03.411041Z","iopub.status.idle":"2025-03-22T15:46:07.113673Z","shell.execute_reply.started":"2025-03-22T15:46:03.411015Z","shell.execute_reply":"2025-03-22T15:46:07.113010Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# BioPython을 사용한 FASTA/MSA 파일 로드 함수\ndef load_msa_with_biopython(fasta_file, convert_to_rna=False):  # 기본값을 False로 변경\n    \"\"\"\n    BioPython의 SeqIO를 사용하여 FASTA 형식의 MSA 파일에서 시퀀스를 로드합니다.\n    \n    Args:\n        fasta_file (str): FASTA 파일 경로\n        convert_to_rna (bool): DNA를 RNA로 변환할지 여부 (T→U)\n        \n    Returns:\n        list: 시퀀스 리스트 (첫 번째는 쿼리 시퀀스)\n    \"\"\"\n    try:\n        from Bio import SeqIO\n        sequences = []\n        \n        # FASTA 파일에서 레코드 읽기\n        for record in SeqIO.parse(fasta_file, \"fasta\"):\n            # 시퀀스를 문자열로 변환\n            seq_str = str(record.seq).upper()\n            \n            # DNA를 RNA로 변환 (필요한 경우)\n            if convert_to_rna:\n                seq_str = seq_str.replace('T', 'U')\n            \n            # 소문자 처리 (일부 MSA 형식에서는 삽입을 나타냄)\n            # 여기서는 모든 문자를 대문자로 유지\n            \n            sequences.append(seq_str)\n        \n        # 시퀀스가 없는 경우 기본값 반환\n        if not sequences:\n            return [\"\"]\n        \n        return sequences\n    \n    except ImportError:\n        print(\"BioPython이 설치되어 있지 않습니다. 기본 파서를 사용합니다.\")\n        return load_a3m_fallback(fasta_file, convert_to_rna)\n    except Exception as e:\n        print(f\"FASTA 파일 로드 중 오류 발생: {str(e)}\")\n        return [\"\"]\n\n# BioPython 없을 경우를 위한 대체 파서\ndef load_a3m_fallback(fasta_file, convert_to_rna=False):  # 기본값을 False로 변경\n    \"\"\"기본 파서를 사용하여 FASTA 파일을 로드합니다.\"\"\"\n    sequences = []\n    current_seq = \"\"\n    \n    with open(fasta_file, \"r\") as f:\n        for line in f:\n            line = line.strip()\n            if line.startswith(\">\"):\n                if current_seq:\n                    sequences.append(current_seq)\n                    current_seq = \"\"\n            else:\n                # 대문자로 정규화\n                clean_line = line.upper()\n                # DNA를 RNA로 변환 (필요한 경우)\n                if convert_to_rna:\n                    clean_line = clean_line.replace('T', 'U')\n                \n                current_seq += clean_line\n    \n    if current_seq:\n        sequences.append(current_seq)\n    \n    # 시퀀스가 없는 경우를 대비\n    if not sequences:\n        return [\"\"]\n    \n    return sequences\n\n# MSA 데이터를 처리하는 함수\ndef process_msa(msa_sequences, max_seq=32):\n    \"\"\"MSA 시퀀스를 처리하여 숫자 인코딩 텐서로 변환합니다.\"\"\"\n    # 시퀀스가 없거나 빈 문자열인 경우 처리\n    if not msa_sequences or not msa_sequences[0]:\n        # 기본값으로 1x1 배열 반환 (나중에 처리 가능하도록)\n        return np.zeros((1, 1), dtype=np.int64)\n    \n    # 첫 번째 시퀀스는 쿼리 시퀀스\n    query_seq = msa_sequences[0]\n    \n    # 최대 시퀀스 수 제한\n    msa_sequences = msa_sequences[:max_seq]\n    \n    # 시퀀스 길이 확인\n    seq_len = len(query_seq)\n    \n    # 시퀀스가 비어있으면 기본값 반환\n    if seq_len == 0:\n        return np.zeros((len(msa_sequences), 1), dtype=np.int64)\n    \n    # 원-핫 인코딩을 위한 사전 (A, C, G, U, T, - 갭, 기타 문자)\n    # 'T'를 추가하여 DNA와 RNA 모두 처리 가능하도록 합니다\n    vocab = {'A': 0, 'C': 1, 'G': 2, 'U': 3, 'T': 4, '-': 5, '.': 5}\n    \n    # MSA 시퀀스를 숫자로 변환\n    msa_encoded = np.zeros((len(msa_sequences), seq_len), dtype=np.int64)\n    \n    # 각 시퀀스 처리\n    for i, seq in enumerate(msa_sequences):\n        # 현재 시퀀스가 쿼리 시퀀스보다 짧으면 패딩\n        if len(seq) < seq_len:\n            seq = seq + '-' * (seq_len - len(seq))\n        \n        # 시퀀스가 쿼리보다 길면 잘라내기\n        seq = seq[:seq_len]\n        \n        # 문자별 처리\n        for j, nt in enumerate(seq):\n            # 소문자 'c'와 'a'는 비공식 삽입(들여쓰기) 표시로 사용될 수 있음\n            # 여기서는 일반 문자로 처리\n            nt_upper = nt.upper()\n            \n            if nt_upper in vocab:\n                msa_encoded[i, j] = vocab[nt_upper]\n            else:\n                # 모르는 문자는 갭으로 처리\n                msa_encoded[i, j] = vocab['-']\n    \n    return msa_encoded\n\n# 데이터 로드 및 처리\ndef load_data(config):\n    # 기본 데이터 로드\n    train_sequences = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_sequences.csv\")\n    train_labels = pd.read_csv(\"/kaggle/input/stanford-rna-3d-folding/train_labels.csv\")\n    train_labels[\"pdb_id\"] = train_labels[\"ID\"].apply(lambda x: x.split(\"_\")[0] + '_' + x.split(\"_\")[1])\n    \n    all_xyz = []\n    \n    for pdb_id in tqdm(train_sequences['target_id']):\n        df = train_labels[train_labels[\"pdb_id\"] == pdb_id]\n        xyz = df[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n        xyz[xyz < -1e17] = float('NaN')\n        all_xyz.append(xyz)\n    \n    # 필터링\n    filter_nan = []\n    max_len = 0\n    for xyz in all_xyz:\n        if len(xyz) > max_len:\n            max_len = len(xyz)\n        \n        filter_nan.append((np.isnan(xyz).mean() <= 0.5) & \n                         (len(xyz) < config['max_len_filter']) & \n                         (len(xyz) > config['min_len_filter']))\n    \n    print(f\"Longest sequence in train: {max_len}\")\n    \n    filter_nan = np.array(filter_nan)\n    non_nan_indices = np.arange(len(filter_nan))[filter_nan]\n    \n    train_sequences = train_sequences.loc[non_nan_indices].reset_index(drop=True)\n    all_xyz = [all_xyz[i] for i in non_nan_indices]\n    \n    # MSA 정보 추가\n    msa_dir = \"/kaggle/input/stanford-rna-3d-folding/MSA\"\n    msa_data = []\n    \n    for target_id in tqdm(train_sequences['target_id']):\n        msa_file = os.path.join(msa_dir, f\"{target_id}.MSA.fasta\")\n        \n        if os.path.exists(msa_file):\n            try:\n                # BioPython을 사용한 MSA 로드\n                msa_sequences = load_msa_with_biopython(msa_file, convert_to_rna=config['convert_to_rna'])\n                \n                # MSA 정보 출력 (디버깅)\n                if len(msa_sequences) > 1:\n                    print(f\"MSA 로드 성공: {target_id}, {len(msa_sequences)} 시퀀스\")\n                \n                msa_data.append(msa_sequences)\n            except Exception as e:\n                print(f\"MSA 로드 실패: {target_id}, 오류: {str(e)}\")\n                # 기본값: 원본 시퀀스만 포함\n                seq = train_sequences.loc[train_sequences['target_id'] == target_id, 'sequence'].values[0]\n                msa_data.append([seq])\n        else:\n            # MSA 파일이 없는 경우, 원래 시퀀스만 포함\n            seq = train_sequences.loc[train_sequences['target_id'] == target_id, 'sequence'].values[0]\n            msa_data.append([seq])\n            \n            if target_id.startswith('8S'):  # 디버깅을 위해 몇 개의 누락된 파일만 기록\n                print(f\"MSA 파일 없음: {target_id}\")\n    \n    # 데이터 패키징\n    data = {\n        \"sequence\": train_sequences['sequence'].to_list(),\n        \"temporal_cutoff\": train_sequences['temporal_cutoff'].to_list(),\n        \"description\": train_sequences['description'].to_list(),\n        \"all_sequences\": train_sequences['all_sequences'].to_list(),\n        \"xyz\": all_xyz,\n        \"msa\": msa_data\n    }\n    \n    # 데이터 분할\n    all_index = np.arange(len(data['sequence']))\n    cutoff_date = pd.Timestamp(config['cutoff_date'])\n    test_cutoff_date = pd.Timestamp(config['test_cutoff_date'])\n    train_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) <= cutoff_date]\n    test_index = [i for i, d in enumerate(data['temporal_cutoff']) if pd.Timestamp(d) > cutoff_date and pd.Timestamp(d) <= test_cutoff_date]\n    \n    print(f\"Train size: {len(train_index)}\")\n    print(f\"Test size: {len(test_index)}\")\n    \n    return data, train_index, test_index\n\n# RNA MSA 데이터셋 클래스\nclass RNA3D_MSA_Dataset(Dataset):\n    def __init__(self, indices, data, config):\n        self.indices = indices\n        self.data = data\n        self.config = config\n        # 'T'를 추가하여 DNA와 RNA 모두 처리 가능하도록 수정\n        self.tokens = {'A': 0, 'C': 1, 'G': 2, 'U': 3, 'T': 4}\n        self.max_seq = config['msa_max_sequences']\n    \n    def __len__(self):\n        return len(self.indices)\n    \n    def __getitem__(self, idx):\n        idx = self.indices[idx]\n        \n        # 시퀀스 처리\n        try:\n            sequence = [self.tokens.get(nt, 0) for nt in self.data['sequence'][idx]]\n            sequence = np.array(sequence)\n            sequence = torch.tensor(sequence)\n        except Exception as e:\n            print(f\"시퀀스 처리 오류: {str(e)}\")\n            # 기본값 설정 - 길이 1의 시퀀스\n            sequence = torch.tensor([0])\n        \n        # MSA 데이터 가져오기 및 처리\n        try:\n            msa_sequences = self.data['msa'][idx]\n            msa_encoded = process_msa(msa_sequences, self.max_seq)\n            msa_encoded = torch.tensor(msa_encoded)\n        except Exception as e:\n            print(f\"MSA 처리 오류: {str(e)}\")\n            # 기본값 설정 - 1x1 크기의 MSA\n            msa_encoded = torch.tensor([[0]])\n        \n        # XYZ 좌표 가져오기\n        try:\n            xyz = self.data['xyz'][idx]\n            xyz = torch.tensor(np.array(xyz))\n        except Exception as e:\n            print(f\"XYZ 처리 오류: {str(e)}\")\n            # 기본값 설정 - 1x3 크기의 좌표\n            xyz = torch.tensor([[0.0, 0.0, 0.0]])\n        \n        # 데이터 차원 일관성 확인\n        seq_len = len(sequence)\n        \n        # 필요한 경우 시퀀스 자르기\n        if seq_len > self.config['max_len']:\n            crop_start = np.random.randint(seq_len - self.config['max_len'])\n            crop_end = crop_start + self.config['max_len']\n            \n            sequence = sequence[crop_start:crop_end]\n            \n            # xyz 배열이 충분히 길면 자르기\n            if len(xyz) >= crop_end:\n                xyz = xyz[crop_start:crop_end]\n            \n            # msa 배열이 2차원이고 두 번째 차원이 충분히 길면 자르기\n            if len(msa_encoded.shape) == 2 and msa_encoded.shape[1] >= crop_end:\n                msa_encoded = msa_encoded[:, crop_start:crop_end]\n        \n        # MSA와 시퀀스 길이가 맞지 않는 경우 조정\n        if len(msa_encoded.shape) == 2 and msa_encoded.shape[1] != len(sequence):\n            # 더 작은 길이로 자르거나 패딩\n            min_len = min(msa_encoded.shape[1], len(sequence))\n            sequence = sequence[:min_len]\n            \n            if msa_encoded.shape[1] > min_len:\n                msa_encoded = msa_encoded[:, :min_len]\n            elif msa_encoded.shape[1] < min_len:\n                # 패딩 추가\n                padding = torch.zeros((msa_encoded.shape[0], min_len - msa_encoded.shape[1]), dtype=torch.long)\n                msa_encoded = torch.cat([msa_encoded, padding], dim=1)\n        \n        return {\n            'sequence': sequence,\n            'xyz': xyz,\n            'msa': msa_encoded\n        }\n\n# MSA Transformer 모델 구현\nclass MSAAttention(nn.Module):\n    def __init__(self, d_model, n_heads, dropout=0.1):\n        super(MSAAttention, self).__init__()\n        self.d_model = d_model\n        self.n_heads = n_heads\n        self.head_dim = d_model // n_heads\n        \n        self.query = nn.Linear(d_model, d_model)\n        self.key = nn.Linear(d_model, d_model)\n        self.value = nn.Linear(d_model, d_model)\n        self.out = nn.Linear(d_model, d_model)\n        \n        self.dropout = nn.Dropout(dropout)\n    \n    def forward(self, x, mask=None):\n        batch_size, num_seqs, seq_len, _ = x.shape\n        \n        q = self.query(x).view(batch_size, num_seqs, seq_len, self.n_heads, self.head_dim)\n        k = self.key(x).view(batch_size, num_seqs, seq_len, self.n_heads, self.head_dim)\n        v = self.value(x).view(batch_size, num_seqs, seq_len, self.n_heads, self.head_dim)\n        \n        # 차원 변경: batch_size, n_heads, num_seqs, seq_len, head_dim\n        q = q.permute(0, 3, 1, 2, 4)\n        k = k.permute(0, 3, 1, 2, 4)\n        v = v.permute(0, 3, 1, 2, 4)\n        \n        # 어텐션 계산\n        scores = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5)\n        \n        if mask is not None:\n            scores = scores.masked_fill(mask == 0, -1e9)\n        \n        attn_weights = F.softmax(scores, dim=-1)\n        attn_weights = self.dropout(attn_weights)\n        \n        output = torch.matmul(attn_weights, v)\n        \n        # 차원 되돌리기\n        output = output.permute(0, 2, 3, 1, 4).contiguous()\n        output = output.view(batch_size, num_seqs, seq_len, self.d_model)\n        \n        return self.out(output)\n\nclass RowAttention(nn.Module):\n    \"\"\"시퀀스 위치 간의 어텐션 (각 MSA 시퀀스 내에서)\"\"\"\n    def __init__(self, d_model, n_heads, dropout=0.1):\n        super(RowAttention, self).__init__()\n        self.attention = MSAAttention(d_model, n_heads, dropout)\n    \n    def forward(self, x, mask=None):\n        return self.attention(x, mask)\n\nclass ColumnAttention(nn.Module):\n    \"\"\"MSA 시퀀스 간의 어텐션 (각 위치에서)\"\"\"\n    def __init__(self, d_model, n_heads, dropout=0.1):\n        super(ColumnAttention, self).__init__()\n        self.attention = MSAAttention(d_model, n_heads, dropout)\n    \n    def forward(self, x, mask=None):\n        batch_size, num_seqs, seq_len, d_model = x.shape\n        \n        # 차원 변경: (batch_size, seq_len, num_seqs, d_model)\n        x = x.permute(0, 2, 1, 3)\n        \n        if mask is not None:\n            mask = mask.permute(0, 2, 1, 3)\n        \n        x = self.attention(x, mask)\n        \n        # 차원 되돌리기\n        x = x.permute(0, 2, 1, 3)\n        \n        return x\n\nclass FeedForward(nn.Module):\n    def __init__(self, d_model, d_ff=2048, dropout=0.1):\n        super(FeedForward, self).__init__()\n        self.linear1 = nn.Linear(d_model, d_ff)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(d_ff, d_model)\n    \n    def forward(self, x):\n        return self.linear2(self.dropout(F.relu(self.linear1(x))))\n\nclass LayerNorm(nn.Module):\n    def __init__(self, d_model, eps=1e-6):\n        super(LayerNorm, self).__init__()\n        self.gamma = nn.Parameter(torch.ones(d_model))\n        self.beta = nn.Parameter(torch.zeros(d_model))\n        self.eps = eps\n    \n    def forward(self, x):\n        mean = x.mean(-1, keepdim=True)\n        std = x.std(-1, keepdim=True)\n        return self.gamma * (x - mean) / (std + self.eps) + self.beta\n\nclass MSATransformerLayer(nn.Module):\n    def __init__(self, d_model, n_heads, d_ff=2048, dropout=0.1):\n        super(MSATransformerLayer, self).__init__()\n        self.row_attn = RowAttention(d_model, n_heads, dropout)\n        self.col_attn = ColumnAttention(d_model, n_heads, dropout)\n        self.feed_forward = FeedForward(d_model, d_ff, dropout)\n        \n        self.norm1 = LayerNorm(d_model)\n        self.norm2 = LayerNorm(d_model)\n        self.norm3 = LayerNorm(d_model)\n        \n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        self.dropout3 = nn.Dropout(dropout)\n    \n    def forward(self, x, mask=None):\n        # 행 어텐션 (시퀀스 위치 간)\n        row_attn_output = x + self.dropout1(self.row_attn(self.norm1(x), mask))\n        \n        # 열 어텐션 (MSA 시퀀스 간)\n        col_attn_output = row_attn_output + self.dropout2(self.col_attn(self.norm2(row_attn_output), mask))\n        \n        # 피드포워드\n        output = col_attn_output + self.dropout3(self.feed_forward(self.norm3(col_attn_output)))\n        \n        return output\n\nclass MSATransformer(nn.Module):\n    def __init__(self, d_model, n_heads, num_layers, d_ff=2048, dropout=0.1):\n        super(MSATransformer, self).__init__()\n        self.layers = nn.ModuleList([\n            MSATransformerLayer(d_model, n_heads, d_ff, dropout)\n            for _ in range(num_layers)\n        ])\n        self.norm = LayerNorm(d_model)\n    \n    def forward(self, x, mask=None):\n        for layer in self.layers:\n            x = layer(x, mask)\n        return self.norm(x)\n\nclass CrossAttention(nn.Module):\n    def __init__(self, d_model, n_heads, dropout=0.1):\n        super(CrossAttention, self).__init__()\n        self.d_model = d_model\n        self.n_heads = n_heads\n        self.head_dim = d_model // n_heads\n        \n        self.query = nn.Linear(d_model, d_model)\n        self.key = nn.Linear(d_model, d_model)\n        self.value = nn.Linear(d_model, d_model)\n        self.out = nn.Linear(d_model, d_model)\n        \n        self.dropout = nn.Dropout(dropout)\n    \n    def forward(self, query, key_value, mask=None):\n        batch_size = query.shape[0]\n        \n        q = self.query(query)\n        k = self.key(key_value)\n        v = self.value(key_value)\n        \n        # 어텐션 차원 변경\n        q = q.view(batch_size, -1, self.n_heads, self.head_dim).transpose(1, 2)\n        k = k.view(batch_size, -1, self.n_heads, self.head_dim).transpose(1, 2)\n        v = v.view(batch_size, -1, self.n_heads, self.head_dim).transpose(1, 2)\n        \n        # 어텐션 계산\n        scores = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5)\n        \n        if mask is not None:\n            scores = scores.masked_fill(mask == 0, -1e9)\n        \n        attn_weights = F.softmax(scores, dim=-1)\n        attn_weights = self.dropout(attn_weights)\n        \n        output = torch.matmul(attn_weights, v)\n        \n        # 차원 되돌리기\n        output = output.transpose(1, 2).contiguous()\n        output = output.view(batch_size, -1, self.d_model)\n        \n        return self.out(output)\n\nclass CrossTransformerLayer(nn.Module):\n    def __init__(self, d_model, n_heads, d_ff=2048, dropout=0.1):\n        super(CrossTransformerLayer, self).__init__()\n        self.cross_attn = CrossAttention(d_model, n_heads, dropout)\n        self.feed_forward = FeedForward(d_model, d_ff, dropout)\n        \n        self.norm1 = LayerNorm(d_model)\n        self.norm2 = LayerNorm(d_model)\n        self.norm3 = LayerNorm(d_model)\n        \n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n    \n    def forward(self, x, msa_repr, mask=None):\n        # 크로스 어텐션\n        cross_attn_output = x + self.dropout1(self.cross_attn(self.norm1(x), self.norm2(msa_repr), mask))\n        \n        # 피드포워드\n        output = cross_attn_output + self.dropout2(self.feed_forward(self.norm3(cross_attn_output)))\n        \n        return output\n\nclass CrossTransformer(nn.Module):\n    def __init__(self, d_model, n_heads, num_layers, d_ff=2048, dropout=0.1):\n        super(CrossTransformer, self).__init__()\n        self.layers = nn.ModuleList([\n            CrossTransformerLayer(d_model, n_heads, d_ff, dropout)\n            for _ in range(num_layers)\n        ])\n        self.norm = LayerNorm(d_model)\n    \n    def forward(self, x, msa_repr, mask=None):\n        for layer in self.layers:\n            x = layer(x, msa_repr, mask)\n        return self.norm(x)\n\nclass StructureModule(nn.Module):\n    def __init__(self, d_model, n_heads, num_layers, d_ff=2048, dropout=0.1):\n        super(StructureModule, self).__init__()\n        self.layers = nn.ModuleList([\n            MSATransformerLayer(d_model, n_heads, d_ff, dropout)\n            for _ in range(num_layers)\n        ])\n        self.norm = LayerNorm(d_model)\n        self.xyz_predictor = nn.Linear(d_model, 3)\n    \n    def forward(self, x, mask=None):\n        for layer in self.layers:\n            x = layer(x, mask)\n        x = self.norm(x)\n        xyz = self.xyz_predictor(x)\n        return xyz\n\nclass RNA_MSA_Folding(nn.Module):\n    def __init__(self, config):\n        super(RNA_MSA_Folding, self).__init__()\n        self.config = config\n        \n        # 임베딩 레이어 - 'T' 토큰 추가로 6개 (A, C, G, U, T, 갭(-))\n        self.seq_embedding = nn.Embedding(6, config['msa_feat_dim'])\n        \n        # 위치 인코딩\n        self.pos_embedding = nn.Parameter(torch.zeros(1, 1, config['max_len'], config['msa_feat_dim']))\n        \n        # MSA Transformer\n        self.msa_transformer = MSATransformer(\n            d_model=config['msa_feat_dim'],\n            n_heads=config['n_heads'],\n            num_layers=config['num_self_attn_layers'],\n            dropout=config['dropout']\n        )\n        \n        # 단일 시퀀스 임베딩을 위한 추가 레이어\n        self.single_seq_embedding = nn.Linear(config['msa_feat_dim'], config['msa_feat_dim'])\n        \n        # Cross Attention\n        self.cross_transformer = CrossTransformer(\n            d_model=config['msa_feat_dim'],\n            n_heads=config['n_heads'],\n            num_layers=config['num_cross_attn_layers'],\n            dropout=config['dropout']\n        )\n        \n        # 구조 모듈\n        self.structure_module = StructureModule(\n            d_model=config['msa_feat_dim'],\n            n_heads=config['n_heads'],\n            num_layers=config['num_structure_module_layers'],\n            dropout=config['dropout']\n        )\n    \n    def forward(self, seq, msa):\n        # 입력 형태 확인 및 조정\n        if len(msa.shape) != 3:\n            print(f\"Warning: MSA 형태가 비정상적입니다. 형태: {msa.shape}\")\n            # 최소 3차원 텐서로 조정\n            if len(msa.shape) == 2:\n                msa = msa.unsqueeze(0)\n            elif len(msa.shape) == 1:\n                msa = msa.unsqueeze(0).unsqueeze(0)\n        \n        batch_size, num_seqs, seq_len = msa.shape\n        \n        # MSA 임베딩\n        msa_emb = self.seq_embedding(msa)  # [batch_size, num_seqs, seq_len, d_model]\n        \n        # 위치 인코딩 추가 (너무 긴 시퀀스 처리)\n        if seq_len <= self.pos_embedding.shape[2]:\n            pos_emb = self.pos_embedding[:, :, :seq_len, :]\n        else:\n            # 위치 인코딩 확장\n            existing_pos = self.pos_embedding.squeeze(0).squeeze(0)  # [max_len, d_model]\n            extended_pos = torch.zeros(seq_len, self.config['msa_feat_dim'], device=msa.device)\n            extended_pos[:existing_pos.shape[0], :] = existing_pos\n            # 나머지 부분은 마지막 위치 인코딩 복제\n            if existing_pos.shape[0] > 0:\n                extended_pos[existing_pos.shape[0]:, :] = existing_pos[-1, :]\n            pos_emb = extended_pos.unsqueeze(0).unsqueeze(0)\n            \n        msa_emb = msa_emb + pos_emb\n        \n        # MSA 처리\n        msa_repr = self.msa_transformer(msa_emb)  # [batch_size, num_seqs, seq_len, d_model]\n        \n        # 쿼리 시퀀스 (첫 번째 시퀀스)에 대한 임베딩\n        query_repr = msa_repr[:, 0]  # [batch_size, seq_len, d_model]\n        query_repr = self.single_seq_embedding(query_repr)\n        \n        # MSA 정보를 쿼리 시퀀스에 크로스 어텐션으로 통합\n        query_refined = self.cross_transformer(query_repr, msa_repr)\n        \n        # 구조 모듈로 좌표 예측\n        xyz_pred = self.structure_module(query_refined.unsqueeze(1)).squeeze(1)\n        \n        return xyz_pred\n\n\n# 단일 시퀀스 버전의 RNA Folding 모델\nclass RNA_Single_Folding(nn.Module):\n    def __init__(self, config):\n        super(RNA_Single_Folding, self).__init__()\n        self.config = config\n        \n        # 임베딩 레이어 - 'T' 토큰 추가로 5개 (A, C, G, U, T)\n        self.seq_embedding = nn.Embedding(5, config['msa_feat_dim'])\n        \n        # 위치 인코딩\n        self.pos_embedding = nn.Parameter(torch.zeros(1, config['max_len'], config['msa_feat_dim']))\n        \n        # 자기주의 트랜스포머 레이어 (MSA 없이 단일 시퀀스에 대해서만)\n        self.transformer_layers = nn.ModuleList([\n            nn.TransformerEncoderLayer(\n                d_model=config['msa_feat_dim'],\n                nhead=config['n_heads'],\n                dim_feedforward=config['msa_feat_dim'] * 4,\n                dropout=config['dropout'],\n                batch_first=True\n            ) for _ in range(config['num_self_attn_layers'] + config['num_cross_attn_layers'])\n        ])\n        \n        self.norm = LayerNorm(config['msa_feat_dim'])\n        \n        # 구조 모듈\n        self.structure_module = StructureModule(\n            d_model=config['msa_feat_dim'],\n            n_heads=config['n_heads'],\n            num_layers=config['num_structure_module_layers'],\n            dropout=config['dropout']\n        )\n    \n    def forward(self, seq, msa=None):  # msa는 인터페이스 호환을 위해 있지만 사용하지 않음\n        # 시퀀스 임베딩\n        seq_emb = self.seq_embedding(seq)  # [batch_size, seq_len, d_model]\n        \n        # 시퀀스 길이 확인\n        seq_len = seq_emb.shape[1]\n        \n        # 위치 인코딩 추가\n        if seq_len <= self.pos_embedding.shape[1]:\n            pos_emb = self.pos_embedding[:, :seq_len, :]\n        else:\n            # 위치 인코딩 확장\n            existing_pos = self.pos_embedding.squeeze(0)  # [max_len, d_model]\n            extended_pos = torch.zeros(seq_len, self.config['msa_feat_dim'], device=seq.device)\n            extended_pos[:existing_pos.shape[0], :] = existing_pos\n            # 나머지 부분은 마지막 위치 인코딩 복제\n            if existing_pos.shape[0] > 0:\n                extended_pos[existing_pos.shape[0]:, :] = existing_pos[-1, :]\n            pos_emb = extended_pos.unsqueeze(0)\n            \n        seq_emb = seq_emb + pos_emb\n        \n        # 트랜스포머 레이어 통과\n        x = seq_emb\n        for layer in self.transformer_layers:\n            x = layer(x)\n        \n        x = self.norm(x)\n        \n        # 구조 모듈로 좌표 예측\n        xyz_pred = self.structure_module(x.unsqueeze(1)).squeeze(1)\n        \n        return xyz_pred\n\n# 거리 행렬 계산 함수\ndef calculate_distance_matrix(X, Y, epsilon=1e-4):\n    return (torch.square(X[:, None] - Y[None, :]) + epsilon).sum(-1).sqrt()\n\n# dRMAE 손실 함수\ndef dRMAE(pred_x, pred_y, gt_x, gt_y, epsilon=1e-4, Z=10, d_clamp=None):\n    pred_dm = calculate_distance_matrix(pred_x, pred_y)\n    gt_dm = calculate_distance_matrix(gt_x, gt_y)\n    \n    mask = ~torch.isnan(gt_dm)\n    mask[torch.eye(mask.shape[0]).bool()] = False\n    \n    rmsd = torch.abs(pred_dm[mask] - gt_dm[mask])\n    \n    return rmsd.mean() / Z\n\n# SVD 기반 정렬 및 MAE 손실 함수\ndef align_svd_mae(input, target, Z=10):\n    \"\"\"\n    SVD 기반 Procrustes 정렬을 사용하여 입력을 타겟에 정렬하고 MAE 손실을 계산합니다.\n    \"\"\"\n    assert input.shape == target.shape, \"입력과 타겟의 형태가 같아야 합니다\"\n    \n    # 마스크 적용\n    mask = ~torch.isnan(target.sum(-1))\n    \n    input = input[mask]\n    target = target[mask]\n    \n    # 중심점 계산\n    centroid_input = input.mean(dim=0, keepdim=True)\n    centroid_target = target.mean(dim=0, keepdim=True)\n    \n    # 중심 정규화\n    input_centered = input - centroid_input.detach()\n    target_centered = target - centroid_target\n    \n    # 공분산 행렬 계산\n    cov_matrix = input_centered.T @ target_centered\n    \n    # SVD를 사용한 최적 회전 찾기\n    U, S, Vt = torch.svd(cov_matrix)\n    \n    # 회전 행렬 계산\n    R = Vt @ U.T\n    \n    # 회전이 적절한지 확인 (det(R) = 1, 반사 없음)\n    if torch.det(R) < 0:\n        Vt[-1, :] *= -1\n        R = Vt @ U.T\n    \n    # 입력 회전\n    aligned_input = (input_centered @ R.T.detach()) + centroid_target.detach()\n    \n    return torch.abs(aligned_input - target).mean() / Z","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:46:07.118603Z","iopub.execute_input":"2025-03-22T15:46:07.118943Z","iopub.status.idle":"2025-03-22T15:46:07.181862Z","shell.execute_reply.started":"2025-03-22T15:46:07.118909Z","shell.execute_reply":"2025-03-22T15:46:07.181063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rigid_transform_3D_safe(A, B):\n    \"\"\"\n    자동 미분과 호환되는 안전한 방식으로 최적의 회전 및 변환 행렬을 계산합니다.\n    \n    Args:\n        A: (N, 3) 형태의 텐서 - 예측 좌표\n        B: (N, 3) 형태의 텐서 - 실제 좌표\n        \n    Returns:\n        R: 회전 행렬 (3, 3)\n        t: 변환 벡터 (3)\n    \"\"\"\n    # NaN 값 제거\n    mask = ~torch.isnan(B).any(dim=1)\n    A_filtered = A[mask]\n    B_filtered = B[mask]\n    \n    if len(A_filtered) < 3:  # 유효한 변환을 위해 최소 3개의 점 필요\n        # 항등 회전 및 0 변환 반환\n        return torch.eye(3, device=A.device), torch.zeros(3, device=A.device)\n    \n    # 중심점 계산\n    centroid_A = torch.mean(A_filtered, dim=0)\n    centroid_B = torch.mean(B_filtered, dim=0)\n    \n    # 중심 정규화\n    A_centered = A_filtered - centroid_A\n    B_centered = B_filtered - centroid_B\n    \n    # 공분산 행렬 계산\n    H = A_centered.T @ B_centered\n    \n    # SVD - 자동 미분과 호환되는 방식으로\n    try:\n        # 표준 SVD\n        U, S, V = torch.linalg.svd(H, full_matrices=False)\n    except Exception:\n        # 문제가 있으면 더 안정적인 방법 시도\n        print(\"표준 SVD 실패, 안정화된 방법 시도\")\n        # 작은 값 추가로 안정화\n        H_stable = H + torch.eye(3, device=H.device) * 1e-6\n        U, S, V = torch.linalg.svd(H_stable, full_matrices=False)\n    \n    # 회전 행렬 계산\n    R = V.T @ U.T\n    \n    # 결정자 확인 (반사가 아닌 회전인지)\n    det = torch.det(R)\n    if det < 0:\n        # 클론을 생성하여 인플레이스 연산 방지\n        V_adjusted = V.clone()\n        V_adjusted[-1] = V_adjusted[-1] * -1\n        R = V_adjusted.T @ U.T\n    \n    # 변환 계산\n    t = centroid_B - R @ centroid_A\n    \n    return R, t\n\ndef fape_loss_safe(pred, target, clamp_distance=10.0, eps=1e-8):\n    \"\"\"\n    자동 미분과 호환되는 안전한 방식으로 Frame Aligned Point Error (FAPE) 손실을 계산합니다.\n    \n    Args:\n        pred: (N, 3) 형태의 예측 좌표 텐서\n        target: (N, 3) 형태의 실제 좌표 텐서\n        clamp_distance: 최대 거리 (이상치 처리용)\n        eps: 0으로 나누기 방지를 위한 작은 값\n        \n    Returns:\n        FAPE 손실 값\n    \"\"\"\n    # target의 NaN 값 처리\n    mask = ~torch.isnan(target).any(dim=1)\n    \n    # 유효한 점이 없으면 0 손실 반환\n    if mask.sum() == 0:\n        return torch.tensor(0.0, device=pred.device, requires_grad=True)\n    \n    # 최적의 회전 및 변환 계산\n    try:\n        R, t = rigid_transform_3D_safe(pred, target)\n        \n        # 예측값을 target에 정렬\n        pred_aligned = (R @ pred.T).T + t\n        \n        # 유효한 위치에서만 오차 계산\n        error = torch.sqrt(torch.sum((pred_aligned[mask] - target[mask])**2, dim=1) + eps)\n        \n        # 큰 거리는 제한하여 이상치 영향 감소\n        error = torch.clamp(error, max=clamp_distance)\n        \n        return error.mean()\n    except Exception as e:\n        print(f\"FAPE 손실 계산 중 오류: {str(e)}\")\n        # 오류 발생 시 대체 손실 반환\n        return torch.mean(torch.abs(pred[mask] - target[mask]))\n\ndef combined_rna_loss_safe(pred_coords, target_coords, sequence, \n                         fape_weight=1.0,\n                         rmsd_weight=0.5,\n                         violation_weight=0.2,\n                         structural_violation_epoch=50,\n                         current_epoch=0):\n    \"\"\"\n    RNA 구조 예측을 위한 안전한 조합 손실 함수\n    \n    Args:\n        pred_coords: (N, 3) 형태의 예측 좌표 텐서\n        target_coords: (N, 3) 형태의 실제 좌표 텐서\n        sequence: RNA 시퀀스 인덱스\n        fape_weight: FAPE 손실 가중치\n        rmsd_weight: RMSD 손실 가중치\n        violation_weight: 구조적 위반 가중치\n        structural_violation_epoch: 구조적 위반 적용 시작 에폭\n        current_epoch: 현재 훈련 에폭\n        \n    Returns:\n        조합 손실 값\n    \"\"\"\n    # FAPE 손실 - 안전한 버전 사용\n    fape = fape_loss_safe(pred_coords, target_coords)\n    \n    # RMSD 손실\n    try:\n        rmsd = dRMAE(pred_coords, pred_coords, target_coords, target_coords)\n    except Exception:\n        # 오류 발생 시 대체 손실\n        mask = ~torch.isnan(target_coords).any(dim=1)\n        if mask.sum() > 0:\n            rmsd = torch.mean(torch.abs(pred_coords[mask] - target_coords[mask]))\n        else:\n            rmsd = torch.tensor(0.0, device=pred_coords.device, requires_grad=True)\n    \n    # 총 손실 초기화\n    total_loss = fape_weight * fape + rmsd_weight * rmsd\n    \n    # 특정 에폭 이후 구조적 위반 추가\n    if current_epoch >= structural_violation_epoch:\n        try:\n            # 기존 structural_violations_loss 함수 호출\n            violations = structural_violations_loss(pred_coords, sequence, violation_weight)\n            total_loss += violations\n        except Exception as e:\n            print(f\"구조적 위반 계산 중 오류: {str(e)}\")\n            # 오류 발생 시 구조적 위반 무시\n            pass\n    \n    return total_loss\n\n\n# 4. Curriculum Learning Implementation\n\nclass CurriculumSampler:\n    \"\"\"\n    Implements curriculum learning by gradually increasing the complexity\n    of training examples based on RNA length and other features.\n    \"\"\"\n    def __init__(self, dataset, max_epochs, min_len=10, max_len=500):\n        self.dataset = dataset\n        self.max_epochs = max_epochs\n        self.min_len = min_len\n        self.max_len = max_len\n        self.epoch = 0\n        \n        # Calculate sequence lengths for all examples\n        self.lengths = []\n        for idx in dataset.indices:\n            seq_len = len(dataset.data['sequence'][idx])\n            self.lengths.append(seq_len)\n        \n        self.lengths = np.array(self.lengths)\n        self.indices = np.arange(len(dataset.indices))\n        \n        # Initial sort by length\n        self.sorted_indices = self.indices[np.argsort(self.lengths)]\n        \n    def update_epoch(self, epoch):\n        \"\"\"Update current epoch\"\"\"\n        self.epoch = epoch\n        \n    def get_indices(self):\n        \"\"\"\n        Returns indices for the current epoch based on curriculum.\n        Gradually includes longer sequences as training progresses.\n        \"\"\"\n        # Calculate progress ratio (0 to 1)\n        progress = min(1.0, self.epoch / (self.max_epochs * 0.8))\n        \n        # Calculate max length for current epoch\n        current_max_len = self.min_len + progress * (self.max_len - self.min_len)\n        \n        # Get indices of sequences up to current max length\n        mask = self.lengths <= current_max_len\n        curriculum_indices = self.indices[mask]\n        \n        # If too few sequences, include at least 20% of the data\n        if len(curriculum_indices) < len(self.indices) * 0.2:\n            min_count = int(len(self.indices) * 0.2)\n            curriculum_indices = self.sorted_indices[:min_count]\n        \n        # Shuffle the selected indices\n        np.random.shuffle(curriculum_indices)\n        \n        return curriculum_indices","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:48:11.052896Z","iopub.execute_input":"2025-03-22T15:48:11.053216Z","iopub.status.idle":"2025-03-22T15:48:11.069576Z","shell.execute_reply.started":"2025-03-22T15:48:11.053192Z","shell.execute_reply":"2025-03-22T15:48:11.068713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 훈련 함수\ndef train_model(model, train_loader, val_loader, config, model_path=None):\n    # 저장된 모델 불러오기 (있는 경우)\n    if model_path and os.path.exists(model_path):\n        print(f\"모델 불러오기: {model_path}\")\n        try:\n            model.load_state_dict(torch.load(model_path))\n            print(\"모델 로드 성공!\")\n        except Exception as e:\n            print(f\"모델 로드 실패: {str(e)}\")\n            print(\"새로운 모델로 훈련을 시작합니다.\")\n    else:\n        print(\"사전 훈련된 모델이 없습니다. 새로운 모델로 훈련을 시작합니다.\")\n    \n    optimizer = torch.optim.Adam(model.parameters(), lr=config['learning_rate'], weight_decay=config['weight_decay'])\n    \n    scaler = GradScaler()\n    \n    # 코사인 스케줄러\n    schedule = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, \n        T_max=(config['epochs'] - config['cos_epoch']) * len(train_loader) // config['batch_size']\n    )\n    \n    best_val_loss = float('inf')\n    \n    # 이슈 추적을 위한 디버그 모드 (필요한 경우 활성화)\n    # torch.autograd.set_detect_anomaly(True)\n    \n    for epoch in range(config['epochs']):\n        model.train()\n        tbar = tqdm(train_loader)\n        total_loss = 0\n        \n        for idx, batch in enumerate(tbar):\n            sequence = batch['sequence'].cuda()\n            msa = batch['msa'].cuda()\n            gt_xyz = batch['xyz'].cuda()\n            \n            pred_xyz = model(sequence, msa)\n            \n            # 배치 차원 제거\n            pred_xyz = pred_xyz.squeeze(0)\n            gt_xyz = gt_xyz.squeeze(0)\n            \n            # 새로운 안전한 손실 함수 사용\n            if epoch < config.get('structural_violation_epoch', 50):\n                # 초기에는 기존 손실 함수로 훈련 (안정성을 위해)\n                loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz) + align_svd_mae(pred_xyz, gt_xyz)\n            else:\n                # 일정 에폭 이후 개선된 손실 함수 사용\n                try:\n                    loss = combined_rna_loss_safe(\n                        pred_xyz, \n                        gt_xyz, \n                        sequence,\n                        fape_weight=1.0,\n                        rmsd_weight=0.5,\n                        violation_weight=0.2,\n                        structural_violation_epoch=config.get('structural_violation_epoch', 50),\n                        current_epoch=epoch\n                    )\n                except Exception as e:\n                    print(f\"손실 계산 중 오류: {str(e)}\")\n                    # 오류 발생 시 기존 손실 함수 사용\n                    loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz) + align_svd_mae(pred_xyz, gt_xyz)\n            \n            (loss / config['batch_size']).backward()\n            \n            if (idx + 1) % config['batch_size'] == 0 or idx + 1 == len(tbar):\n                torch.nn.utils.clip_grad_norm_(model.parameters(), config['grad_clip'])\n                optimizer.step()\n                optimizer.zero_grad()\n                \n                if (epoch + 1) > config['cos_epoch']:\n                    schedule.step()\n            \n            total_loss += loss.item()\n            \n            tbar.set_description(f\"Epoch {epoch + 1} Loss: {total_loss / (idx + 1)}\")\n        \n        # 검증\n        model.eval()\n        val_loss = 0\n        val_preds = []\n        \n        tbar = tqdm(val_loader)\n        for idx, batch in enumerate(tbar):\n            sequence = batch['sequence'].cuda()\n            msa = batch['msa'].cuda()\n            gt_xyz = batch['xyz'].cuda()\n            \n            with torch.no_grad():\n                pred_xyz = model(sequence, msa)\n                \n                # 배치 차원 제거\n                pred_xyz = pred_xyz.squeeze(0)\n                gt_xyz = gt_xyz.squeeze(0)\n                \n                # 검증에는 기존 손실 함수 사용 (안정성을 위해)\n                loss = dRMAE(pred_xyz, pred_xyz, gt_xyz, gt_xyz)\n                \n                # 에폭 50 이후에는 추가 검증 메트릭 출력\n                if epoch >= config.get('structural_violation_epoch', 50):\n                    try:\n                        fape = fape_loss_safe(pred_xyz, gt_xyz)\n                        print(f\"배치 {idx} FAPE: {fape.item():.4f}\")\n                    except:\n                        pass\n            \n            val_loss += loss.item()\n            val_preds.append([gt_xyz.cpu().numpy(), pred_xyz.cpu().numpy()])\n        \n        val_loss = val_loss / len(tbar)\n        print(f\"Validation loss: {val_loss}\")\n        \n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            best_preds = val_preds\n            torch.save(model.state_dict(), 'RNA_MSA_Folding_best.pt')\n    \n    # 최종 모델 저장\n    torch.save(model.state_dict(), 'RNA_MSA_Folding_final.pt')\n    \n    return best_val_loss, best_preds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:48:23.513924Z","iopub.execute_input":"2025-03-22T15:48:23.514206Z","iopub.status.idle":"2025-03-22T15:48:23.526245Z","shell.execute_reply.started":"2025-03-22T15:48:23.514183Z","shell.execute_reply":"2025-03-22T15:48:23.525418Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config.update({\n    \"structural_violation_epoch\": 25,  # 25번째 에폭부터 구조 위반 손실 적용\n    \"mixed_precision\": \"bf16\",         # 혼합 정밀도 사용\n    \"gradient_accumulation_steps\": 4,  # 효과적인 배치 크기 증가를 위한 그래디언트 누적\n})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:46:07.229051Z","iopub.execute_input":"2025-03-22T15:46:07.229346Z","iopub.status.idle":"2025-03-22T15:46:07.246303Z","shell.execute_reply.started":"2025-03-22T15:46:07.229323Z","shell.execute_reply":"2025-03-22T15:46:07.245552Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    print(\"데이터 로드 중...\")\n    data, train_index, test_index = load_data(config)\n    \n    print(f\"MSA 데이터 로드 완료: {len(data['msa'])} 항목\")\n    print(f\"첫 번째 MSA 시퀀스 샘플: {data['msa'][0][:2] if data['msa'] and len(data['msa'][0]) > 1 else 'N/A'}\")\n    \n    # 데이터셋 및 데이터로더 생성\n    print(\"데이터셋 생성 중...\")\n    train_dataset = RNA3D_MSA_Dataset(train_index, data, config)\n    val_dataset = RNA3D_MSA_Dataset(test_index, data, config)\n    \n    # 데이터셋 샘플 확인\n    sample = train_dataset[0]\n    print(f\"데이터셋 샘플 형태:\")\n    print(f\"  시퀀스: {sample['sequence'].shape}\")\n    print(f\"  MSA: {sample['msa'].shape}\")\n    print(f\"  XYZ: {sample['xyz'].shape}\")\n    \n    train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\n    val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\n\nexcept Exception as e:\n    print(f\"오류 발생: {str(e)}\")\n    import traceback\n    traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:46:07.248519Z","iopub.execute_input":"2025-03-22T15:46:07.248763Z","iopub.status.idle":"2025-03-22T15:46:31.600018Z","shell.execute_reply.started":"2025-03-22T15:46:07.248741Z","shell.execute_reply":"2025-03-22T15:46:31.599105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predictions(pred_data):\n    \"\"\"\n    예측된 구조와 실제 구조를 시각화합니다.\n    \"\"\"\n    import plotly.graph_objects as go\n    \n    gt_xyz, pred_xyz = pred_data\n    \n    # NaN 값 필터링\n    mask = ~np.isnan(gt_xyz).any(axis=1)\n    gt_xyz_filtered = gt_xyz[mask]\n    pred_xyz_filtered = pred_xyz[mask]\n    \n    # 실제 구조\n    fig1 = go.Figure(data=[go.Scatter3d(\n        x=gt_xyz_filtered[:, 0],\n        y=gt_xyz_filtered[:, 1],\n        z=gt_xyz_filtered[:, 2],\n        mode='markers',\n        marker=dict(\n            size=5,\n            color=np.arange(len(gt_xyz_filtered)),\n            colorscale='Viridis',\n            opacity=0.8\n        ),\n        name='Ground Truth'\n    )])\n    \n    fig1.update_layout(\n        title=\"Ground Truth RNA 3D Structure\",\n        scene=dict(\n            xaxis_title=\"X\",\n            yaxis_title=\"Y\",\n            zaxis_title=\"Z\"\n        )\n    )\n    \n    # 예측 구조\n    fig2 = go.Figure(data=[go.Scatter3d(\n        x=pred_xyz_filtered[:, 0],\n        y=pred_xyz_filtered[:, 1],\n        z=pred_xyz_filtered[:, 2],\n        mode='markers',\n        marker=dict(\n            size=5,\n            color=np.arange(len(pred_xyz_filtered)),\n            colorscale='Viridis',\n            opacity=0.8\n        ),\n        name='Prediction'\n    )])\n    \n    fig2.update_layout(\n        title=\"Predicted RNA 3D Structure\",\n        scene=dict(\n            xaxis_title=\"X\",\n            yaxis_title=\"Y\",\n            zaxis_title=\"Z\"\n        )\n    )\n    \n    fig1.show()\n    fig2.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:48:27.526902Z","iopub.execute_input":"2025-03-22T15:48:27.527187Z","iopub.status.idle":"2025-03-22T15:48:27.533538Z","shell.execute_reply.started":"2025-03-22T15:48:27.527165Z","shell.execute_reply":"2025-03-22T15:48:27.532806Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# MSA 모델 생성 및 학습\nprint(\"MSA 모델 학습 시작...\")\nmodel_msa = RNA_MSA_Folding(config).cuda()\n\n# 사전 훈련된 모델 경로\npretrained_model_path = \"/kaggle/input/rna-msa-folding/RNA_MSA_Folding_model.pt\"\nbest_val_loss_msa, best_preds_msa = train_model(model_msa, train_loader, val_loader, config, pretrained_model_path)\n\n# 단일 시퀀스 모델 생성 및 학습 (옵션)\nprint(\"단일 시퀀스 모델 학습 시작...\")\nmodel_single = RNA_Single_Folding(config).cuda()\nbest_val_loss_single, best_preds_single = train_model(model_single, train_loader, val_loader, config)\n\n# 결과 비교\nprint(f\"MSA 모델 최종 검증 손실: {best_val_loss_msa}\")\nprint(f\"단일 시퀀스 모델 최종 검증 손실: {best_val_loss_single}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:48:27.766903Z","iopub.execute_input":"2025-03-22T15:48:27.767188Z","execution_failed":"2025-03-22T17:51:51.853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 예측 시각화\nif best_preds:\n    print(\"예측 결과 시각화 중...\")\n    visualize_predictions(best_preds[0])\nelse:\n    print(\"시각화를 위한 예측 결과가 없습니다.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T15:46:38.023280Z","iopub.status.idle":"2025-03-22T15:46:38.023527Z","shell.execute_reply":"2025-03-22T15:46:38.023427Z"}},"outputs":[],"execution_count":null}]}