{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false}],"dockerImageVersionId":31287,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"5f3e0a46","cell_type":"code","source":"# ============================================================\n# Stanford RNA 3D Folding Part 2 - V5.1 CPU estável, eficiente e balanceado\n# ------------------------------------------------------------\n# Ajustes principais:\n# - logs com timestamp e sem duplicação\n# - parser robusto para train/validation labels\n# - positional encoding dinâmica (corrige sequências > 4096)\n# - fallback seguro de validação\n# - mensagem de erro mais clara por etapa\n# - submission.csv válido\n# ============================================================\n\nimport os\nimport json\nimport hashlib\nimport math\nimport time\nimport random\nimport logging\nimport warnings\nimport traceback\nimport pickle\n\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\n\nwarnings.filterwarnings(\"ignore\")\n\n# ============================================================\n# 1) LOGGING\n# ============================================================\nLOGGER_NAME = \"rna3d_v42\"\nlogger = logging.getLogger(LOGGER_NAME)\nlogger.setLevel(logging.INFO)\nlogger.propagate = False\nlogger.handlers.clear()\n\nhandler = logging.StreamHandler()\nformatter = logging.Formatter(\n    fmt=\"%(asctime)s | %(levelname)-7s | %(message)s\",\n    datefmt=\"%H:%M:%S\"\n)\nhandler.setFormatter(formatter)\nlogger.addHandler(handler)\n\ndef log_section(title: str):\n    logger.info(\"=\" * 72)\n    logger.info(title)\n    logger.info(\"=\" * 72)\n\ndef log_df_info(name: str, df: pd.DataFrame, max_cols: int = 12):\n    logger.info(\"%s shape=%s\", name, df.shape)\n    logger.info(\"%s columns=%s\", name, list(df.columns[:max_cols]))\n\n# ============================================================\n# 2) CONFIG\n# ============================================================\nSEED = 42\nVERSION_TAG = \"V5.3_AUTHORAL_REFINED_COSINE\"\nMAX_EPOCHS = 7\nBATCH_SIZE = 2\nLR = 5.0e-4\nMIN_LR = 8.0e-5\nWEIGHT_DECAY = 1.0e-2\nD_MODEL = 128\nNHEAD = 4\nNUM_LAYERS = 3\nDROPOUT = 0.14\nMC_SAMPLES = 10\nGRAD_CLIP = 0.45\nPATIENCE = 3\nPE_BASE_MAX_LEN = 4096\nTRAIN_MAX_LEN = 1152\nTRAIN_MID_MAX_LEN = 1664\nINFER_CHUNK_LEN = 896\nINFER_CHUNK_OVERLAP = 192\nMAX_TRAIN_SAMPLES = 5000\nCACHE_DIR = \"/kaggle/working/rna3d_cache_v53\"\nUSE_SAMPLE_CACHE = True\nDIST_LOSS_MAX_POINTS = 112\nVAL_MAX_SAMPLES = 112\nPAIR_SEARCH_MAX_SPAN = 256\nPAIR_SEARCH_TOPK = 6\nLENGTH_BINS_FOR_CAP = 10\nPAIR_LOSS_W = 0.02\nSMOOTH_LOSS_W = 0.20\nSTEP_LOSS_W = 0.10\nSTERIC_LOSS_W_BASE = 0.04\nSTERIC_LOSS_W_FINAL = 0.12\nCURRENT_EPOCH_FRACTION = 0.0\nTRAIN_LENGTH_HARD_CAP = 4096\nWARMUP_EPOCHS = 1\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\nos.makedirs(CACHE_DIR, exist_ok=True)\n\nlogger.info(\"DEVICE=%s\", DEVICE)\nlogger.info(\"SEED=%s\", SEED)\n\n# ============================================================\n# 3) FILE DISCOVERY\n# ============================================================\ndef find_file(filename, root=\"/kaggle/input\"):\n    for r, _, files in os.walk(root):\n        if filename in files:\n            return os.path.join(r, filename)\n    return None\n\ndef find_first_existing(filenames, root=\"/kaggle/input\"):\n    for name in filenames:\n        p = find_file(name, root=root)\n        if p is not None:\n            return p\n    return None\n\ndef print_version_banner():\n    line = \"=\" * 72\n    logger.info(line)\n    logger.info(\"VERSION_TAG=%s\", VERSION_TAG)\n    logger.info(\"RNA 3D Folding CPU notebook com loss por distâncias e sem Kabsch na loss principal\")\n    logger.info(line)\n\n# ============================================================\n# 4) HELPERS DE SCHEMA\n# ============================================================\ndef find_col(df, candidates):\n    cols_lower = {str(c).lower(): c for c in df.columns}\n    for cand in candidates:\n        if cand.lower() in cols_lower:\n            return cols_lower[cand.lower()]\n    for c in df.columns:\n        cl = str(c).lower()\n        for cand in candidates:\n            if cand.lower() in cl:\n                return c\n    return None\n\ndef normalize_sequence_df(df, df_name=\"sequence_df\"):\n    log_section(f\"NORMALIZE {df_name}\")\n    id_col = find_col(df, [\"ID\", \"target_id\", \"target\", \"sequence_id\", \"chain\"])\n    seq_col = find_col(df, [\"sequence\", \"seq\", \"rna_sequence\", \"resname\", \"resnames\"])\n\n    logger.info(\"%s id_col=%s\", df_name, id_col)\n    logger.info(\"%s seq_col=%s\", df_name, seq_col)\n\n    if id_col is None:\n        raise ValueError(f\"Não encontrei coluna de ID em {df_name}: {list(df.columns)}\")\n    if seq_col is None:\n        raise ValueError(f\"Não encontrei coluna de sequência em {df_name}: {list(df.columns)}\")\n\n    out = df[[id_col, seq_col]].copy()\n    out.columns = [\"target_id\", \"sequence\"]\n    out[\"target_id\"] = out[\"target_id\"].astype(str)\n    out[\"sequence\"] = out[\"sequence\"].astype(str).str.upper().str.replace(\" \", \"\", regex=False)\n    logger.info(\"%s normalized shape=%s\", df_name, out.shape)\n    return out\n\ndef normalize_label_df(df, df_name=\"label_df\"):\n    \"\"\"\n    Aceita:\n    1) formato longo: ID,resname,resid,x,y,z\n    2) formato longo com x_1,y_1,z_1\n    3) formato largo por alvo com x_1..x_n, y_1..y_n, z_1..z_n\n    \"\"\"\n    log_section(f\"NORMALIZE {df_name}\")\n\n    # -------- Caso 1: formato longo --------\n    id_col = find_col(df, [\"ID\"])\n    resname_col = find_col(df, [\"resname\", \"base\", \"nucleotide\"])\n    resid_col = find_col(df, [\"resid\", \"residue\", \"position\", \"idx\", \"index\"])\n    x_col = find_col(df, [\"x\", \"x_1\"])\n    y_col = find_col(df, [\"y\", \"y_1\"])\n    z_col = find_col(df, [\"z\", \"z_1\"])\n\n    logger.info(\n        \"%s long-format candidates -> ID=%s resname=%s resid=%s x=%s y=%s z=%s\",\n        df_name, id_col, resname_col, resid_col, x_col, y_col, z_col\n    )\n\n    if None not in [id_col, resname_col, resid_col, x_col, y_col, z_col]:\n        out = df[[id_col, resname_col, resid_col, x_col, y_col, z_col]].copy()\n        out.columns = [\"ID\", \"resname\", \"resid\", \"x\", \"y\", \"z\"]\n        out[\"ID\"] = out[\"ID\"].astype(str)\n        out[\"resname\"] = out[\"resname\"].astype(str).str.upper()\n        out[\"resid\"] = pd.to_numeric(out[\"resid\"], errors=\"coerce\")\n        out[\"x\"] = pd.to_numeric(out[\"x\"], errors=\"coerce\")\n        out[\"y\"] = pd.to_numeric(out[\"y\"], errors=\"coerce\")\n        out[\"z\"] = pd.to_numeric(out[\"z\"], errors=\"coerce\")\n        out = out.dropna(subset=[\"resid\", \"x\", \"y\", \"z\"]).copy()\n        out[\"resid\"] = out[\"resid\"].astype(int)\n        out[\"target_id\"] = out[\"ID\"].apply(lambda s: \"_\".join(str(s).split(\"_\")[:-1]))\n        logger.info(\"%s detected LONG label format\", df_name)\n        logger.info(\"%s normalized shape=%s\", df_name, out.shape)\n        return out[[\"ID\", \"resname\", \"resid\", \"x\", \"y\", \"z\", \"target_id\"]]\n\n    # -------- Caso 2: formato largo --------\n    target_col = find_col(df, [\"target_id\", \"ID\", \"id\", \"target\", \"sequence_id\", \"chain\"])\n    seq_col = find_col(df, [\"sequence\", \"seq\", \"rna_sequence\", \"resname\", \"resnames\"])\n\n    x_cols = [c for c in df.columns if str(c).lower().startswith(\"x_\")]\n    y_cols = [c for c in df.columns if str(c).lower().startswith(\"y_\")]\n    z_cols = [c for c in df.columns if str(c).lower().startswith(\"z_\")]\n\n    def sort_xyz(cols):\n        def key_fn(s):\n            try:\n                return int(str(s).split(\"_\")[-1])\n            except Exception:\n                return 10**9\n        return sorted(cols, key=key_fn)\n\n    x_cols = sort_xyz(x_cols)\n    y_cols = sort_xyz(y_cols)\n    z_cols = sort_xyz(z_cols)\n\n    logger.info(\n        \"%s wide-format candidates -> target_col=%s seq_col=%s x=%d y=%d z=%d\",\n        df_name, target_col, seq_col, len(x_cols), len(y_cols), len(z_cols)\n    )\n\n    if target_col is not None and len(x_cols) > 0 and len(x_cols) == len(y_cols) == len(z_cols):\n        rows = []\n        for _, row in df.iterrows():\n            tid = str(row[target_col])\n            seq = str(row[seq_col]).strip().upper() if seq_col is not None else None\n\n            for i in range(len(x_cols)):\n                x = pd.to_numeric(row[x_cols[i]], errors=\"coerce\")\n                y = pd.to_numeric(row[y_cols[i]], errors=\"coerce\")\n                z = pd.to_numeric(row[z_cols[i]], errors=\"coerce\")\n                if pd.isna(x) or pd.isna(y) or pd.isna(z):\n                    continue\n\n                base = seq[i] if seq is not None and i < len(seq) else \"A\"\n                rows.append({\n                    \"ID\": f\"{tid}_{i+1}\",\n                    \"resname\": base,\n                    \"resid\": i + 1,\n                    \"x\": float(x),\n                    \"y\": float(y),\n                    \"z\": float(z),\n                    \"target_id\": tid,\n                })\n\n        out = pd.DataFrame(rows)\n        if len(out) > 0:\n            logger.info(\"%s detected WIDE label format\", df_name)\n            logger.info(\"%s normalized shape=%s\", df_name, out.shape)\n            return out[[\"ID\", \"resname\", \"resid\", \"x\", \"y\", \"z\", \"target_id\"]]\n\n    raise ValueError(f\"{df_name}: schema inesperado. Primeiras colunas: {list(df.columns)[:30]}\")\n\n# ============================================================\n# 5) FEATURES\n# ============================================================\nBASES = [\"A\", \"C\", \"G\", \"U\", \"T\", \"N\"]\nBASE2IDX = {b: i for i, b in enumerate(BASES)}\n\nPAIR_SCORE = {\n    (\"A\", \"U\"): 1.0, (\"U\", \"A\"): 1.0,\n    (\"A\", \"T\"): 1.0, (\"T\", \"A\"): 1.0,\n    (\"G\", \"C\"): 1.2, (\"C\", \"G\"): 1.2,\n    (\"G\", \"U\"): 0.7, (\"U\", \"G\"): 0.7,\n    (\"G\", \"T\"): 0.7, (\"T\", \"G\"): 0.7,\n}\n\ndef clean_base(b):\n    b = str(b).upper()\n    return b if b in BASE2IDX else \"N\"\n\ndef seq_to_ids(seq):\n    return np.array([BASE2IDX[clean_base(b)] for b in seq], dtype=np.int64)\n\n\n\ndef heuristic_pair_features(seq, min_loop=4, max_span=PAIR_SEARCH_MAX_SPAN, topk=PAIR_SEARCH_TOPK):\n    L = len(seq)\n    pair_map = np.full(L, -1, dtype=np.int64)\n    pair_prob = np.zeros(L, dtype=np.float32)\n    pair_type = np.zeros(L, dtype=np.int64)  # 0=unpaired, 1=AU/UA, 2=CG/GC, 3=GU/UG\n    used = np.zeros(L, dtype=bool)\n    neighbor_cands = [[] for _ in range(L)]\n\n    # Busca local-limitada: bem mais eficiente que o O(L^2) total e suficiente para capturar stems.\n    for i in range(L):\n        lo = i + min_loop + 1\n        hi = min(L, i + max_span + 1)\n        bi = seq[i]\n        if lo >= hi:\n            continue\n        for j in range(lo, hi):\n            bj = seq[j]\n            pair = (bi, bj)\n            base_sc = PAIR_SCORE.get(pair, 0.0)\n            if base_sc <= 0.0:\n                continue\n\n            span = j - i\n            span_bonus = 1.0 / (1.0 + abs(span - 14) / 20.0)\n\n            stem_bonus = 0.0\n            if i + 1 < j - 1 and PAIR_SCORE.get((seq[i + 1], seq[j - 1]), 0.0) > 0:\n                stem_bonus += 0.22\n            if i - 1 >= 0 and j + 1 < L and PAIR_SCORE.get((seq[i - 1], seq[j + 1]), 0.0) > 0:\n                stem_bonus += 0.18\n\n            gc_bonus = 0.10 if pair in [(\"G\", \"C\"), (\"C\", \"G\")] else 0.0\n            long_penalty = 0.92 if span > (max_span * 0.75) else 1.0\n            score = (base_sc * span_bonus + stem_bonus + gc_bonus) * long_penalty\n\n            neighbor_cands[i].append((score, i, j, pair))\n            neighbor_cands[j].append((score, i, j, pair))\n\n    cands = []\n    for local in neighbor_cands:\n        if local:\n            local.sort(key=lambda x: x[0], reverse=True)\n            cands.extend(local[:topk])\n\n    uniq = {}\n    for score, i, j, pair in cands:\n        key = (min(i, j), max(i, j))\n        if key not in uniq or score > uniq[key][0]:\n            uniq[key] = (score, i, j, pair)\n\n    cands = list(uniq.values())\n    cands.sort(key=lambda x: x[0], reverse=True)\n\n    for score, i, j, pair in cands:\n        if used[i] or used[j]:\n            continue\n        used[i] = True\n        used[j] = True\n        pair_map[i] = j\n        pair_map[j] = i\n        prob = min(1.0, max(0.0, score / 1.6))\n        pair_prob[i] = prob\n        pair_prob[j] = prob\n        if pair in [(\"A\", \"U\"), (\"U\", \"A\")]:\n            t = 1\n        elif pair in [(\"C\", \"G\"), (\"G\", \"C\")]:\n            t = 2\n        else:\n            t = 3\n        pair_type[i] = t\n        pair_type[j] = t\n\n    paired_flag = (pair_map >= 0).astype(np.float32)\n    return pair_map, pair_prob, pair_type, paired_flag\n\ndef heuristic_pair_map(seq, min_loop=4):\n    pair_map, _, _, _ = heuristic_pair_features(seq, min_loop=min_loop)\n    return pair_map\n\n# ============================================================\n# 6) MONTAR SAMPLES\n# ============================================================\ndef build_samples(seq_df, lbl_df=None, name=\"samples\"):\n    logger.info(\"Building %s...\", name)\n    seq_map = dict(zip(seq_df[\"target_id\"], seq_df[\"sequence\"]))\n    samples = []\n\n    if lbl_df is None:\n        for tid, seq in seq_map.items():\n            samples.append({\n                \"target_id\": tid,\n                \"sequence\": seq,\n                \"coords\": None,\n                \"resids\": np.arange(1, len(seq) + 1, dtype=np.int64),\n            })\n        logger.info(\"%s built without labels -> %d samples\", name, len(samples))\n        return samples\n\n    for tid, g in lbl_df.groupby(\"target_id\", sort=False):\n        if tid not in seq_map:\n            continue\n\n        g = g.sort_values(\"resid\").copy()\n        seq = str(seq_map[tid]).upper()\n        coords = g[[\"x\", \"y\", \"z\"]].values.astype(np.float32)\n\n        L = min(len(seq), len(coords))\n        if L < 3:\n            continue\n\n        coords = coords[:L]\n        valid_mask = np.isfinite(coords).all(axis=1).astype(np.float32)\n        coords = np.nan_to_num(coords, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32)\n\n        if valid_mask.sum() < 3:\n            continue\n\n        samples.append({\n            \"target_id\": tid,\n            \"sequence\": seq[:L],\n            \"coords\": coords,\n            \"valid_mask\": valid_mask,\n            \"resids\": np.arange(1, L + 1, dtype=np.int64),\n        })\n\n    logger.info(\"%s built -> %d samples\", name, len(samples))\n    return samples\n\n\ndef cap_samples_diverse_by_length(samples, cap, seed=SEED, n_bins=LENGTH_BINS_FOR_CAP):\n    if len(samples) <= cap:\n        return samples\n\n    lengths = np.array([len(s[\"sequence\"]) for s in samples], dtype=np.int32)\n    order = np.argsort(lengths)\n    bins = np.array_split(order, max(1, min(n_bins, len(samples))))\n    rng = random.Random(seed)\n\n    chosen = []\n    seen = set()\n    per_bin = max(1, cap // max(1, len(bins)))\n\n    for b in bins:\n        idxs = list(map(int, b.tolist()))\n        rng.shuffle(idxs)\n        take = idxs[:per_bin]\n        for idx in take:\n            if idx not in seen:\n                seen.add(idx)\n                chosen.append(samples[idx])\n\n    if len(chosen) < cap:\n        remaining = [int(i) for i in order.tolist() if int(i) not in seen]\n        rng.shuffle(remaining)\n        for idx in remaining[:cap - len(chosen)]:\n            chosen.append(samples[idx])\n\n    return chosen[:cap]\n\n\ndef order_samples_for_training(samples, seed=SEED, bucket_size=64):\n    if len(samples) <= 1:\n        return samples\n    rng = random.Random(seed)\n    samples = sorted(samples, key=lambda x: len(x[\"sequence\"]))\n    ordered = []\n    for i in range(0, len(samples), bucket_size):\n        chunk = samples[i:i + bucket_size]\n        rng.shuffle(chunk)\n        ordered.extend(chunk)\n    return ordered\n\n# ============================================================\n# 7) NORMALIZAÇÃO DE COORDENADAS\n# ============================================================\ndef normalize_coords(coords, valid_mask=None):\n    coords = np.nan_to_num(coords, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32)\n    if valid_mask is None:\n        valid_mask = np.ones(len(coords), dtype=np.float32)\n    vm = valid_mask.astype(bool)\n    if vm.sum() < 3:\n        return coords.astype(np.float32)\n    center = coords[vm].mean(axis=0, keepdims=True)\n    c = coords - center\n    scale = np.sqrt((c[vm] ** 2).sum(axis=1).mean()) + 1e-6\n    c = c / scale\n    c[~vm] = 0.0\n    return c.astype(np.float32)\n\n# ============================================================\n# 7.5) HELPERS DE ESTABILIDADE/CPU\n# ============================================================\ndef finite_tensor(x: torch.Tensor) -> torch.Tensor:\n    return torch.nan_to_num(x, nan=0.0, posinf=1e4, neginf=-1e4)\n\ndef clamp_coords_torch(x: torch.Tensor, limit: float = 25.0) -> torch.Tensor:\n    return torch.clamp(finite_tensor(x), -limit, limit)\n\ndef maybe_crop_sample(sample, max_len=TRAIN_MAX_LEN, mid_len=TRAIN_MID_MAX_LEN):\n    seq = sample[\"sequence\"]\n    L = len(seq)\n    if L <= mid_len:\n        return sample\n    crop_len = max_len if L > 2048 else mid_len\n    crop_len = min(crop_len, L)\n    if crop_len >= L:\n        return sample\n    start = random.randint(0, L - crop_len)\n    end = start + crop_len\n    out = dict(sample)\n    out[\"sequence\"] = sample[\"sequence\"][start:end]\n    out[\"resids\"] = sample[\"resids\"][start:end]\n    if sample.get(\"coords\") is not None:\n        out[\"coords\"] = sample[\"coords\"][start:end]\n        if sample.get(\"valid_mask\") is not None:\n            out[\"valid_mask\"] = sample[\"valid_mask\"][start:end]\n    return out\n\n# ============================================================\n# 8) DATASET / COLLATE\n# ============================================================\n\ndef sample_cache_path(seq_df, lbl_df, name):\n    seq_sig = f\"{len(seq_df)}_{int(seq_df['target_id'].astype(str).str.len().sum())}_{int(seq_df['sequence'].astype(str).str.len().sum())}\"\n    lbl_sig = \"none\"\n    if lbl_df is not None:\n        lbl_sig = f\"{len(lbl_df)}_{int(lbl_df['target_id'].astype(str).str.len().sum())}_{int(pd.to_numeric(lbl_df['resid'], errors='coerce').fillna(0).sum())}\"\n    key = f\"{name}_{seq_sig}_{lbl_sig}\"\n    key = hashlib.md5(key.encode(\"utf-8\")).hexdigest()[:16]\n    return os.path.join(CACHE_DIR, f\"{name}_{key}.pkl\")\n\ndef load_or_build_samples(seq_df, lbl_df=None, name=\"samples\"):\n    cache_path = sample_cache_path(seq_df, lbl_df, name)\n    if USE_SAMPLE_CACHE and os.path.exists(cache_path):\n        try:\n            with open(cache_path, \"rb\") as f:\n                samples = pickle.load(f)\n            logger.info(\"%s loaded from cache -> %d samples\", name, len(samples))\n            return samples\n        except Exception as e:\n            logger.warning(\"Falha ao ler cache %s: %s\", cache_path, e)\n    samples = build_samples(seq_df, lbl_df, name=name)\n    if USE_SAMPLE_CACHE:\n        try:\n            with open(cache_path, \"wb\") as f:\n                pickle.dump(samples, f, protocol=pickle.HIGHEST_PROTOCOL)\n            logger.info(\"%s cached at %s\", name, cache_path)\n        except Exception as e:\n            logger.warning(\"Falha ao salvar cache %s: %s\", cache_path, e)\n    return samples\n\nclass RNADataset(Dataset):\n    def __init__(self, samples, training=False, max_len=TRAIN_MAX_LEN, mid_len=TRAIN_MID_MAX_LEN):\n        self.samples = samples\n        self.training = training\n        self.max_len = max_len\n        self.mid_len = mid_len\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        s = self.samples[idx]\n        if self.training and s.get(\"coords\") is not None:\n            s = maybe_crop_sample(s, max_len=self.max_len, mid_len=self.mid_len)\n\n        seq = [clean_base(b) for b in s[\"sequence\"]]\n        seq_ids = seq_to_ids(seq)\n        L = len(seq_ids)\n\n        pair_map, pair_prob, pair_type, paired_flag = heuristic_pair_features(seq)\n\n        out = {\n            \"target_id\": s[\"target_id\"],\n            \"sequence\": seq,\n            \"seq_ids\": seq_ids,\n            \"pair_map\": pair_map,\n            \"pair_prob\": pair_prob,\n            \"pair_type\": pair_type,\n            \"paired_flag\": paired_flag,\n            \"pos_norm\": (np.arange(L, dtype=np.float32) / max(L - 1, 1)),\n            \"length\": L,\n            \"resids\": s[\"resids\"],\n        }\n\n        if s.get(\"coords\") is not None:\n            valid_mask = s.get(\"valid_mask\", np.ones(L, dtype=np.float32))\n            out[\"valid_mask\"] = valid_mask.astype(np.float32)\n            out[\"coords\"] = normalize_coords(s[\"coords\"], valid_mask=valid_mask)\n\n        return out\n\ndef collate_fn(batch):\n    B = len(batch)\n    Lmax = max(x[\"length\"] for x in batch)\n\n    seq_ids = torch.full((B, Lmax), BASE2IDX[\"N\"], dtype=torch.long)\n    pair_map = torch.full((B, Lmax), -1, dtype=torch.long)\n    pair_prob = torch.zeros((B, Lmax), dtype=torch.float32)\n    pair_type = torch.zeros((B, Lmax), dtype=torch.long)\n    paired_flag = torch.zeros((B, Lmax), dtype=torch.float32)\n    pos_norm = torch.zeros((B, Lmax), dtype=torch.float32)\n    mask = torch.zeros((B, Lmax), dtype=torch.bool)\n\n    has_coords = \"coords\" in batch[0]\n    coords = torch.zeros((B, Lmax, 3), dtype=torch.float32) if has_coords else None\n    valid_mask = torch.zeros((B, Lmax), dtype=torch.float32) if has_coords else None\n\n    meta = []\n    for i, x in enumerate(batch):\n        L = x[\"length\"]\n        seq_ids[i, :L] = torch.from_numpy(x[\"seq_ids\"])\n        pair_map[i, :L] = torch.from_numpy(x[\"pair_map\"])\n        pair_prob[i, :L] = torch.from_numpy(x[\"pair_prob\"])\n        pair_type[i, :L] = torch.from_numpy(x[\"pair_type\"])\n        paired_flag[i, :L] = torch.from_numpy(x[\"paired_flag\"])\n        pos_norm[i, :L] = torch.from_numpy(x[\"pos_norm\"])\n        mask[i, :L] = True\n\n        if has_coords:\n            coords[i, :L] = torch.from_numpy(x[\"coords\"])\n            valid_mask[i, :L] = torch.from_numpy(x.get(\"valid_mask\", np.ones(L, dtype=np.float32)))\n\n        meta.append({\n            \"target_id\": x[\"target_id\"],\n            \"sequence\": x[\"sequence\"],\n            \"resids\": x[\"resids\"],\n            \"length\": x[\"length\"],\n        })\n\n    out = {\n        \"seq_ids\": seq_ids,\n        \"pair_map\": pair_map,\n        \"pair_prob\": pair_prob,\n        \"pair_type\": pair_type,\n        \"paired_flag\": paired_flag,\n        \"pos_norm\": pos_norm,\n        \"mask\": mask,\n        \"meta\": meta,\n    }\n    if has_coords:\n        out[\"coords\"] = coords\n        out[\"valid_mask\"] = valid_mask\n    return out\n\n# ============================================================\n# 9) MODELO\n# ============================================================\nclass SinusoidalPositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=PE_BASE_MAX_LEN):\n        super().__init__()\n        self.d_model = d_model\n        pe = self._build_pe(max_len)\n        self.register_buffer(\"pe\", pe, persistent=False)\n\n    def _build_pe(self, max_len):\n        pe = torch.zeros(max_len, self.d_model)\n        pos = torch.arange(0, max_len, dtype=torch.float32).unsqueeze(1)\n        div = torch.exp(torch.arange(0, self.d_model, 2).float() * (-math.log(10000.0) / self.d_model))\n        pe[:, 0::2] = torch.sin(pos * div)\n        pe[:, 1::2] = torch.cos(pos * div)\n        return pe.unsqueeze(0)\n\n    def _ensure_len(self, needed_len, device):\n        current_len = self.pe.size(1)\n        if needed_len <= current_len:\n            return\n        new_len = current_len\n        while new_len < needed_len:\n            new_len *= 2\n        logger.warning(\"Expanding positional encoding from %d to %d\", current_len, new_len)\n        self.pe = self._build_pe(new_len).to(device)\n\n    def forward(self, x):\n        self._ensure_len(x.size(1), x.device)\n        return x + self.pe[:, :x.size(1)]\n\n\nclass RNA3DNet(nn.Module):\n    def __init__(self, vocab_size, d_model=128, nhead=8, num_layers=3, dropout=0.15):\n        super().__init__()\n        self.embed = nn.Embedding(vocab_size, d_model)\n        self.pos_mlp = nn.Sequential(\n            nn.Linear(1, d_model // 4),\n            nn.GELU(),\n            nn.Linear(d_model // 4, d_model),\n        )\n        self.pair_offset_embed = nn.Embedding(2048, d_model)\n        self.pair_type_embed = nn.Embedding(4, d_model)\n        self.pair_prob_mlp = nn.Sequential(\n            nn.Linear(1, d_model // 4),\n            nn.GELU(),\n            nn.Linear(d_model // 4, d_model),\n        )\n        self.pe = SinusoidalPositionalEncoding(d_model)\n        self.relpos_bias = nn.Embedding(32, nhead)\n\n        self.local_conv = nn.Sequential(\n            nn.Conv1d(d_model, d_model, kernel_size=5, padding=2, groups=1),\n            nn.GELU(),\n            nn.Conv1d(d_model, d_model, kernel_size=3, padding=1, groups=1),\n            nn.GELU(),\n        )\n\n        self.bigru = nn.GRU(\n            input_size=d_model,\n            hidden_size=d_model // 2,\n            num_layers=1,\n            bidirectional=True,\n            batch_first=True\n        )\n\n        enc_layer = nn.TransformerEncoderLayer(\n            d_model=d_model,\n            nhead=nhead,\n            dim_feedforward=d_model * 3,\n            dropout=dropout,\n            activation=\"gelu\",\n            batch_first=True,\n            norm_first=False,\n        )\n        self.encoder = nn.TransformerEncoder(enc_layer, num_layers=num_layers)\n\n        self.coord_head = nn.Sequential(\n            nn.LayerNorm(d_model),\n            nn.Linear(d_model, d_model),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(d_model, 3)\n        )\n        self.pair_head = nn.Sequential(\n            nn.LayerNorm(d_model),\n            nn.Linear(d_model, d_model // 2),\n            nn.GELU(),\n            nn.Linear(d_model // 2, 1)\n        )\n        self.step_head = nn.Sequential(\n            nn.LayerNorm(d_model),\n            nn.Linear(d_model, d_model // 2),\n            nn.GELU(),\n            nn.Linear(d_model // 2, 1),\n            nn.Softplus()\n        )\n\n    def build_relpos_mask(self, L, device, dtype):\n        pos = torch.arange(L, device=device)\n        dist = (pos[:, None] - pos[None, :]).abs()\n        buckets = torch.clamp((torch.log2(dist.float() + 1.0) * 4.0).long(), max=31)\n        bias = self.relpos_bias(buckets).permute(2, 0, 1).mean(dim=0)\n        local_bonus = torch.where(dist <= 32, torch.zeros_like(dist, dtype=dtype), -0.015 * torch.log1p(dist.float()))\n        return (bias.to(dtype) + local_bonus.to(dtype))\n\n    def forward(self, seq_ids, pair_map, pair_prob, pair_type, pos_norm, mask):\n        B, L = seq_ids.shape\n        x = self.embed(seq_ids)\n        x = x + self.pos_mlp(pos_norm.unsqueeze(-1))\n\n        idx = torch.arange(L, device=seq_ids.device).unsqueeze(0).expand(B, L)\n        rel = torch.where(pair_map >= 0, (pair_map - idx).abs(), torch.zeros_like(pair_map))\n        rel = rel.clamp(max=2047)\n\n        x = x + 0.20 * self.pair_offset_embed(rel)\n        x = x + 0.15 * self.pair_type_embed(pair_type.clamp(min=0, max=3))\n        x = x + 0.20 * self.pair_prob_mlp(pair_prob.unsqueeze(-1))\n\n        x = self.pe(x)\n\n        xc = self.local_conv(x.transpose(1, 2)).transpose(1, 2)\n        x = x + 0.30 * xc\n\n        x, _ = self.bigru(x)\n        relpos_mask = self.build_relpos_mask(L, x.device, x.dtype)\n        x = self.encoder(x, mask=relpos_mask, src_key_padding_mask=~mask)\n\n        coords = torch.tanh(self.coord_head(x)) * 4.0\n        pair_logits = self.pair_head(x).squeeze(-1)\n        step_pred = self.step_head(x).squeeze(-1)\n        return {\n            \"coords\": coords,\n            \"pair_logits\": pair_logits,\n            \"step_pred\": step_pred,\n        }\n\n# ============================================================\n# 10) LOSS\n# ============================================================\n\n\ndef kabsch_align_torch(P, Q, mask):\n    # Mantido apenas por compatibilidade; a loss principal da V4.9 não depende mais de SVD/Kabsch.\n    return clamp_coords_torch(P)\n\n\ndef _subsample_valid_points(p, q, max_points=DIST_LOSS_MAX_POINTS):\n    n = p.shape[0]\n    if n <= max_points:\n        return p, q\n    idx = torch.linspace(0, n - 1, steps=max_points, device=p.device).round().long().unique(sorted=True)\n    return p.index_select(0, idx), q.index_select(0, idx)\n\n\ndef pairwise_distance_loss(pred, true, coord_mask, max_points=DIST_LOSS_MAX_POINTS):\n    losses = []\n    for b in range(pred.shape[0]):\n        m = coord_mask[b]\n        if m.sum() < 4:\n            continue\n        p = pred[b, m]\n        q = true[b, m]\n        if (not torch.isfinite(p).all()) or (not torch.isfinite(q).all()):\n            continue\n        p, q = _subsample_valid_points(p, q, max_points=max_points)\n        if p.shape[0] < 4:\n            continue\n        dp = torch.cdist(clamp_coords_torch(p, limit=20.0), clamp_coords_torch(p, limit=20.0))\n        dq = torch.cdist(clamp_coords_torch(q, limit=20.0), clamp_coords_torch(q, limit=20.0))\n        dp = torch.nan_to_num(dp, nan=50.0, posinf=50.0, neginf=0.0)\n        dq = torch.nan_to_num(dq, nan=50.0, posinf=50.0, neginf=0.0)\n        iu = torch.triu_indices(dp.size(0), dp.size(1), offset=1, device=pred.device)\n        if iu.numel() == 0:\n            continue\n        losses.append(nn.functional.smooth_l1_loss(dp[iu[0], iu[1]], dq[iu[0], iu[1]], beta=0.5))\n    if not losses:\n        return pred.new_tensor(0.0)\n    return torch.stack(losses).mean()\n\n\ndef loss_fn(outputs, true, mask, paired_flag, valid_mask=None):\n    pred = clamp_coords_torch(outputs[\"coords\"])\n    true = clamp_coords_torch(true)\n    if valid_mask is None:\n        valid_mask = mask.float()\n    valid_mask = torch.nan_to_num(valid_mask.float(), nan=0.0, posinf=0.0, neginf=0.0).clamp(0.0, 1.0)\n    coord_mask = mask & (valid_mask > 0.5)\n    pair_logits = torch.nan_to_num(outputs[\"pair_logits\"], nan=0.0, posinf=20.0, neginf=-20.0).clamp(-20, 20)\n    step_pred = torch.nan_to_num(outputs[\"step_pred\"], nan=0.0, posinf=10.0, neginf=0.0).clamp(0, 10)\n\n    # Loss principal invariante a rotação/translação sem SVD/Kabsch.\n    dist_loss = pairwise_distance_loss(pred, true, coord_mask)\n\n    m2 = coord_mask[:, 1:] & coord_mask[:, :-1]\n    dp = clamp_coords_torch(pred[:, 1:] - pred[:, :-1], limit=10.0)\n    dt = clamp_coords_torch(true[:, 1:] - true[:, :-1], limit=10.0)\n\n    smooth = nn.functional.smooth_l1_loss(dp, dt, reduction=\"none\", beta=0.5).sum(dim=-1)\n    smooth = (smooth * m2.float()).sum() / (m2.float().sum() + 1e-6)\n\n    step_true = torch.linalg.norm(dt, dim=-1).clamp(0, 10)\n    step_aux = (step_pred[:, :-1] - step_true).abs()\n    step_aux = (step_aux * m2.float()).sum() / (m2.float().sum() + 1e-6)\n\n    pair_loss = nn.functional.binary_cross_entropy_with_logits(pair_logits, paired_flag, reduction=\"none\")\n    pair_loss = (pair_loss * coord_mask.float()).sum() / (coord_mask.float().sum() + 1e-6)\n\n    steric = pred.new_tensor(0.0)\n    steric_n = 0\n    for b in range(pred.shape[0]):\n        m = coord_mask[b]\n        p = pred[b, m]\n        if p.shape[0] < 6:\n            continue\n        p, _dummy = _subsample_valid_points(p, p, max_points=80)\n        p = clamp_coords_torch(p, limit=20.0)\n        dmat = torch.cdist(p, p)\n        dmat = torch.nan_to_num(dmat, nan=99.0, posinf=99.0, neginf=0.0)\n        iu = torch.triu_indices(dmat.size(0), dmat.size(1), offset=3, device=pred.device)\n        vals = dmat[iu[0], iu[1]]\n        if vals.numel() == 0:\n            continue\n        steric = steric + torch.relu(0.90 - vals).mean()\n        steric_n += 1\n    if steric_n > 0:\n        steric = steric / steric_n\n\n    steric_w = STERIC_LOSS_W_BASE + (STERIC_LOSS_W_FINAL - STERIC_LOSS_W_BASE) * float(CURRENT_EPOCH_FRACTION)\n    total = dist_loss + SMOOTH_LOSS_W * smooth + PAIR_LOSS_W * pair_loss + STEP_LOSS_W * step_aux + steric_w * steric\n    total = torch.nan_to_num(total, nan=10.0, posinf=10.0, neginf=10.0)\n    metrics = {\n        \"dist\": float(torch.nan_to_num(dist_loss.detach().cpu(), nan=0.0)),\n        \"smooth\": float(torch.nan_to_num(smooth.detach().cpu(), nan=0.0)),\n        \"pair\": float(torch.nan_to_num(pair_loss.detach().cpu(), nan=0.0)),\n        \"step\": float(torch.nan_to_num(step_aux.detach().cpu(), nan=0.0)),\n        \"steric\": float(torch.nan_to_num(steric.detach().cpu(), nan=0.0)),\n        \"steric_w\": float(steric_w),\n    }\n    return total, metrics\n\n\n# ============================================================\n# 11) LOOP DE TREINO\n# ============================================================\ndef run_epoch(loader, model, optimizer=None):\n    training = optimizer is not None\n    model.train(training)\n    total_loss = 0.0\n    total_n = 0\n    t0 = time.time()\n\n    for step, batch in enumerate(loader, start=1):\n        seq_ids = batch[\"seq_ids\"].to(DEVICE)\n        pair_map = batch[\"pair_map\"].to(DEVICE)\n        pair_prob = batch[\"pair_prob\"].to(DEVICE)\n        pair_type = batch[\"pair_type\"].to(DEVICE)\n        paired_flag = batch[\"paired_flag\"].to(DEVICE)\n        pos_norm = batch[\"pos_norm\"].to(DEVICE)\n        mask = batch[\"mask\"].to(DEVICE)\n        coords = batch[\"coords\"].to(DEVICE)\n        valid_mask = batch.get(\"valid_mask\")\n        valid_mask = valid_mask.to(DEVICE) if valid_mask is not None else None\n\n        with torch.set_grad_enabled(training):\n            outputs = model(seq_ids, pair_map, pair_prob, pair_type, pos_norm, mask)\n            try:\n                loss, loss_parts = loss_fn(outputs, coords, mask, paired_flag, valid_mask=valid_mask)\n            except RuntimeError as e:\n                logger.warning(\"Skipping batch with loss failure at step=%d: %s\", step, e)\n                if training:\n                    optimizer.zero_grad(set_to_none=True)\n                continue\n\n            if not torch.isfinite(loss):\n                logger.warning(\"Skipping non-finite batch at step=%d\", step)\n                continue\n\n            if training:\n                optimizer.zero_grad(set_to_none=True)\n                loss.backward()\n                nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                has_bad_grad = False\n                for p in model.parameters():\n                    if p.grad is not None and (not torch.isfinite(p.grad).all()):\n                        has_bad_grad = True\n                        break\n                if has_bad_grad:\n                    logger.warning(\"Skipping optimizer step with non-finite gradients at step=%d\", step)\n                    optimizer.zero_grad(set_to_none=True)\n                    continue\n                optimizer.step()\n\n        bs = seq_ids.size(0)\n        total_loss += loss.item() * bs\n        total_n += bs\n\n        if step % 50 == 0 or step == len(loader):\n            logger.info(\n                \"%s step=%d/%d batch_loss=%.5f avg_loss=%.5f\",\n                \"TRAIN\" if training else \"VALID\",\n                step, len(loader),\n                loss.item(),\n                total_loss / max(total_n, 1)\n            )\n            logger.info(\"loss_parts dist=%.4f smooth=%.4f pair=%.4f step=%.4f steric=%.4f steric_w=%.3f\", loss_parts[\"dist\"], loss_parts[\"smooth\"], loss_parts[\"pair\"], loss_parts[\"step\"], loss_parts[\"steric\"], loss_parts.get(\"steric_w\", 0.0))\n\n    return total_loss / max(total_n, 1), time.time() - t0\n\n\n# ============================================================\n# 12) INFERÊNCIA COM RERANKING\n# ============================================================\nINFER_DROPOUT_P = 0.24\nDIVERSITY_NOISE_BASE = 0.018\n\ndef enable_mc_dropout(m):\n    if isinstance(m, nn.Dropout):\n        m.train()\n        if hasattr(m, \"p\"):\n            m.p = max(float(m.p), INFER_DROPOUT_P)\n\ndef add_controlled_diversity(pred, sample_idx=0, total_samples=1):\n    pred = np.asarray(pred, dtype=np.float32).copy()\n    L = len(pred)\n    if L < 2:\n        return pred\n    strength = DIVERSITY_NOISE_BASE * (1.0 + (sample_idx / max(total_samples - 1, 1)))\n    idx = np.linspace(0.0, 1.0, L, dtype=np.float32)\n    sinus = np.stack([\n        np.sin((sample_idx + 1) * math.pi * idx),\n        np.cos((sample_idx + 1) * math.pi * idx),\n        np.sin((sample_idx + 1) * 2.0 * math.pi * idx),\n    ], axis=1).astype(np.float32)\n    pred += strength * sinus\n    noise = np.random.normal(0.0, strength * 0.35, size=pred.shape).astype(np.float32)\n    pred += noise\n    pred -= pred.mean(axis=0, keepdims=True)\n    return pred\n\ndef diversify_if_collapsed(raw_preds, target_k=5):\n    if len(raw_preds) <= 1:\n        return raw_preds\n    dists = []\n    for i in range(len(raw_preds)):\n        for j in range(i + 1, len(raw_preds)):\n            dists.append(pairwise_candidate_distance(raw_preds[i], raw_preds[j]))\n    mean_dist = float(np.mean(dists)) if dists else 0.0\n    if mean_dist > 0.025:\n        return raw_preds\n    diversified = []\n    for i, pred in enumerate(raw_preds):\n        diversified.append(add_controlled_diversity(pred, i, max(len(raw_preds), target_k)))\n    while len(diversified) < max(target_k, 5):\n        base = raw_preds[len(diversified) % len(raw_preds)]\n        diversified.append(add_controlled_diversity(base, len(diversified), max(len(raw_preds), target_k)))\n    return diversified\n\ndef smooth_coords(arr, window=5):\n    if len(arr) < window:\n        return arr\n    out = arr.copy()\n    pad = window // 2\n    for c in range(3):\n        x = arr[:, c]\n        xp = np.pad(x, (pad, pad), mode=\"edge\")\n        kernel = np.ones(window, dtype=np.float32) / window\n        out[:, c] = np.convolve(xp, kernel, mode=\"valid\")\n    return out\n\ndef adaptive_window(L):\n    if L < 80:\n        return 3\n    elif L < 300:\n        return 5\n    elif L < 1000:\n        return 7\n    return 9\n\ndef robust_center(preds):\n    stack = np.stack(preds, axis=0)\n    return np.median(stack, axis=0)\n\ndef local_smoothness_score(pred):\n    if len(pred) < 3:\n        return 0.0\n    d1 = pred[1:] - pred[:-1]\n    d2 = d1[1:] - d1[:-1]\n    return float(np.mean(np.linalg.norm(d2, axis=1)))\n\ndef step_length_score(pred):\n    if len(pred) < 2:\n        return 0.0\n    d = np.linalg.norm(pred[1:] - pred[:-1], axis=1)\n    return float(np.std(d))\n\ndef clash_penalty(pred, min_dist=1.25, skip_near=2):\n    L = len(pred)\n    if L < 5:\n        return 0.0\n    penalty = 0.0\n    count = 0\n    for i in range(L):\n        j0 = i + skip_near + 1\n        if j0 >= L:\n            continue\n        diff = pred[j0:] - pred[i]\n        dist = np.linalg.norm(diff, axis=1)\n        bad = np.maximum(0.0, min_dist - dist)\n        penalty += float(np.sum(bad))\n        count += len(dist)\n    return penalty / max(count, 1)\n\ndef compactness_score(pred):\n    rg = np.sqrt(np.mean(np.sum((pred - pred.mean(axis=0, keepdims=True)) ** 2, axis=1)))\n    return float(rg)\n\ndef consensus_score(pred, center):\n    return float(np.mean(np.linalg.norm(pred - center, axis=1)))\n\ndef normalize_feature_list(values, reverse=False):\n    arr = np.asarray(values, dtype=np.float32)\n    if len(arr) == 0:\n        return arr\n    mn = float(arr.min())\n    mx = float(arr.max())\n    if mx - mn < 1e-8:\n        out = np.zeros_like(arr)\n    else:\n        out = (arr - mn) / (mx - mn)\n    if reverse:\n        out = 1.0 - out\n    return out\n\ndef pairwise_candidate_distance(a, b):\n    return float(np.mean(np.linalg.norm(a - b, axis=1)))\n\ndef rank_candidates(cands):\n    center = robust_center([c[\"pred\"] for c in cands])\n    consensus_vals = [consensus_score(c[\"pred\"], center) for c in cands]\n    smooth_vals = [local_smoothness_score(c[\"pred\"]) for c in cands]\n    step_vals = [step_length_score(c[\"pred\"]) for c in cands]\n    clash_vals = [clash_penalty(c[\"pred\"]) for c in cands]\n    compact_vals = [compactness_score(c[\"pred\"]) for c in cands]\n\n    consensus_n = normalize_feature_list(consensus_vals, reverse=True)\n    smooth_n = normalize_feature_list(smooth_vals, reverse=True)\n    step_n = normalize_feature_list(step_vals, reverse=True)\n    clash_n = normalize_feature_list(clash_vals, reverse=True)\n\n    compact_arr = np.asarray(compact_vals, dtype=np.float32)\n    compact_mid = np.median(compact_arr)\n    compact_dev = np.abs(compact_arr - compact_mid)\n    compact_n = normalize_feature_list(compact_dev, reverse=True)\n\n    for i, c in enumerate(cands):\n        score = (\n            0.30 * consensus_n[i] +\n            0.18 * smooth_n[i] +\n            0.16 * step_n[i] +\n            0.26 * clash_n[i] +\n            0.10 * compact_n[i]\n        )\n        c[\"score\"] = float(score)\n        c[\"features\"] = {\n            \"consensus\": float(consensus_vals[i]),\n            \"smooth\": float(smooth_vals[i]),\n            \"step_std\": float(step_vals[i]),\n            \"clash\": float(clash_vals[i]),\n            \"compact\": float(compact_vals[i]),\n        }\n    cands.sort(key=lambda x: x[\"score\"], reverse=True)\n    return cands\n\ndef select_diverse_top_k(cands, k=5, min_div_frac=0.14):\n    if not cands:\n        return []\n    L = len(cands[0][\"pred\"])\n    scale = max(1.5, np.sqrt(L) * min_div_frac)\n\n    selected = [cands[0]]\n    for cand in cands[1:]:\n        ok = True\n        for s in selected:\n            if pairwise_candidate_distance(cand[\"pred\"], s[\"pred\"]) < scale:\n                ok = False\n                break\n        if ok:\n            selected.append(cand)\n        if len(selected) == k:\n            break\n\n    if len(selected) < k:\n        used = {id(x) for x in selected}\n        for cand in cands:\n            if id(cand) not in used:\n                selected.append(cand)\n            if len(selected) == k:\n                break\n    if len(selected) < k and cands:\n        base_preds = [x[\"pred\"] for x in selected] if selected else []\n        filler_idx = 0\n        while len(selected) < k:\n            base = cands[filler_idx % len(cands)]\n            cloned = {\"pred\": add_controlled_diversity(base[\"pred\"], filler_idx + len(selected), k), \"score\": float(base.get(\"score\", 0.0)), \"features\": dict(base.get(\"features\", {}))}\n            selected.append(cloned)\n            filler_idx += 1\n    return selected[:k]\n\ndef sample_predictions_ranked(model, batch, n_samples=MC_SAMPLES, out_k=5):\n    seq_ids = batch[\"seq_ids\"].to(DEVICE)\n    pair_map = batch[\"pair_map\"].to(DEVICE)\n    pair_prob = batch[\"pair_prob\"].to(DEVICE)\n    pair_type = batch[\"pair_type\"].to(DEVICE)\n    pos_norm = batch[\"pos_norm\"].to(DEVICE)\n    mask = batch[\"mask\"].to(DEVICE)\n\n    model.eval()\n    model.apply(enable_mc_dropout)\n\n    raw_preds = []\n    for sidx in range(n_samples):\n        with torch.no_grad():\n            outputs = model(seq_ids, pair_map, pair_prob, pair_type, pos_norm, mask)\n\n        pred = torch.nan_to_num(outputs[\"coords\"][0, mask[0]], nan=0.0, posinf=4.0, neginf=-4.0).detach().cpu().numpy().astype(np.float32)\n        L = pred.shape[0]\n        scale = max(8.0, math.sqrt(L) * 2.25)\n        pred = pred * scale\n        pred = smooth_coords(pred, window=adaptive_window(L))\n        pred = pred - pred.mean(axis=0, keepdims=True)\n        pred = add_controlled_diversity(pred, sidx, n_samples)\n        raw_preds.append(pred)\n\n    raw_preds = diversify_if_collapsed(raw_preds, target_k=max(out_k, n_samples))\n    cands = [{\"pred\": p.copy()} for p in raw_preds]\n    cands = rank_candidates(cands)\n    best = select_diverse_top_k(cands, k=out_k, min_div_frac=0.14)\n    preds = [c[\"pred\"] for c in best]\n\n    while len(preds) < out_k:\n        base = raw_preds[min(len(preds), len(raw_preds) - 1)].copy()\n        preds.append(add_controlled_diversity(base, len(preds), out_k))\n\n    return preds, cands\n\n# ============================================================\n# 13) MAIN\n# ============================================================\ntry:\n    print_version_banner()\n    log_section(\"FILE DISCOVERY\")\n\n    train_seq_path = find_first_existing([\"train_sequences.csv\", \"train.csv\"])\n    train_lbl_path = find_first_existing([\"train_labels.csv\"])\n    val_seq_path   = find_first_existing([\"validation_sequences.csv\", \"validation.csv\"])\n    val_lbl_path   = find_first_existing([\"validation_labels.csv\"])\n    test_seq_path  = find_first_existing([\"test_sequences.csv\", \"test.csv\"])\n    sample_sub_path = find_first_existing([\"sample_submission.csv\"])\n\n    logger.info(\"train_seq_path=%s\", train_seq_path)\n    logger.info(\"train_lbl_path=%s\", train_lbl_path)\n    logger.info(\"val_seq_path  =%s\", val_seq_path)\n    logger.info(\"val_lbl_path  =%s\", val_lbl_path)\n    logger.info(\"test_seq_path =%s\", test_seq_path)\n    logger.info(\"sample_sub_path=%s\", sample_sub_path)\n\n    assert train_seq_path is not None, \"train_sequences.csv/train.csv não encontrado\"\n    assert train_lbl_path is not None, \"train_labels.csv não encontrado\"\n    assert test_seq_path is not None, \"test_sequences.csv/test.csv não encontrado\"\n\n    log_section(\"READ RAW DATA\")\n    train_seq_df = pd.read_csv(train_seq_path, low_memory=False)\n    train_lbl_df = pd.read_csv(train_lbl_path, low_memory=False)\n    val_seq_df = pd.read_csv(val_seq_path, low_memory=False) if val_seq_path else None\n    val_lbl_df = pd.read_csv(val_lbl_path, low_memory=False) if val_lbl_path else None\n    test_seq_df = pd.read_csv(test_seq_path, low_memory=False)\n    sample_sub_df = pd.read_csv(sample_sub_path, low_memory=False) if sample_sub_path else None\n\n    log_df_info(\"train_sequences_raw\", train_seq_df)\n    log_df_info(\"train_labels_raw\", train_lbl_df)\n    log_df_info(\"test_sequences_raw\", test_seq_df)\n    if val_seq_df is not None:\n        log_df_info(\"validation_sequences_raw\", val_seq_df)\n    if val_lbl_df is not None:\n        log_df_info(\"validation_labels_raw\", val_lbl_df)\n\n    log_section(\"NORMALIZE DATAFRAMES\")\n    train_seq_df = normalize_sequence_df(train_seq_df, \"train_sequences\")\n    train_lbl_df = normalize_label_df(train_lbl_df, \"train_labels\")\n    test_seq_df = normalize_sequence_df(test_seq_df, \"test_sequences\")\n\n    try:\n        val_seq_df = normalize_sequence_df(val_seq_df, \"validation_sequences\") if val_seq_df is not None else None\n    except Exception as e:\n        logger.warning(\"Falha ao normalizar validation_sequences: %s\", e)\n        val_seq_df = None\n\n    try:\n        val_lbl_df = normalize_label_df(val_lbl_df, \"validation_labels\") if val_lbl_df is not None else None\n    except Exception as e:\n        logger.warning(\"Falha ao normalizar validation_labels: %s\", e)\n        val_lbl_df = None\n\n    log_df_info(\"train_sequences_norm\", train_seq_df)\n    log_df_info(\"train_labels_norm\", train_lbl_df)\n    log_df_info(\"test_sequences_norm\", test_seq_df)\n    if val_seq_df is not None:\n        log_df_info(\"validation_sequences_norm\", val_seq_df)\n    if val_lbl_df is not None:\n        log_df_info(\"validation_labels_norm\", val_lbl_df)\n\n    if val_seq_df is not None and val_lbl_df is not None:\n        logger.info(\"Mantendo validação separada para monitorar estabilidade e score proxy em CPU.\")\n\n    train_samples = load_or_build_samples(train_seq_df, train_lbl_df, \"train_samples\")\n    if len(train_samples) > MAX_TRAIN_SAMPLES:\n        train_samples = cap_samples_diverse_by_length(train_samples, MAX_TRAIN_SAMPLES, seed=SEED)\n        logger.info(\"train_samples capped for CPU with length balance -> %d\", len(train_samples))\n    over_cap = sum(1 for s in train_samples if len(s[\"sequence\"]) > TRAIN_LENGTH_HARD_CAP)\n    if over_cap:\n        logger.info(\"training samples above hard cap=%d -> %d (will be cropped during training)\", TRAIN_LENGTH_HARD_CAP, over_cap)\n    train_samples = order_samples_for_training(train_samples, seed=SEED)\n    try:\n        val_samples = load_or_build_samples(val_seq_df, val_lbl_df, \"val_samples\") if (val_seq_df is not None and val_lbl_df is not None) else []\n        if len(val_samples) > VAL_MAX_SAMPLES:\n            rng = random.Random(SEED)\n            rng.shuffle(val_samples)\n            val_samples = val_samples[:VAL_MAX_SAMPLES]\n            logger.info(\"val_samples reduced for CPU -> %d\", len(val_samples))\n    except Exception as e:\n        logger.warning(\"Falha ao montar validação, seguindo sem validação: %s\", e)\n        val_samples = []\n    test_samples = load_or_build_samples(test_seq_df, None, \"test_samples\")\n\n    if len(train_samples) > MAX_TRAIN_SAMPLES:\n        train_samples = sorted(train_samples, key=lambda s: len(s[\"sequence\"]))\n        stride = max(1, len(train_samples) // MAX_TRAIN_SAMPLES)\n        train_samples = train_samples[::stride][:MAX_TRAIN_SAMPLES]\n        logger.info(\"train_samples limited for CPU -> %d samples\", len(train_samples))\n\n    train_loader = DataLoader(\n        RNADataset(train_samples, training=True),\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        collate_fn=collate_fn,\n        num_workers=0,\n        pin_memory=(DEVICE == \"cuda\"),\n    )\n    val_loader = DataLoader(\n        RNADataset(val_samples, training=False),\n        batch_size=1,\n        shuffle=False,\n        collate_fn=collate_fn,\n        num_workers=0,\n        pin_memory=(DEVICE == \"cuda\"),\n    ) if len(val_samples) else None\n    test_loader = DataLoader(\n        RNADataset(test_samples, training=False),\n        batch_size=1,\n        shuffle=False,\n        collate_fn=collate_fn,\n        num_workers=0,\n        pin_memory=(DEVICE == \"cuda\"),\n    )\n\n    train_lengths = [len(s[\"sequence\"]) for s in train_samples]\n    logger.info(\n        \"train_lengths stats -> min=%d p50=%d p90=%d max=%d\",\n        int(np.min(train_lengths)), int(np.median(train_lengths)), int(np.percentile(train_lengths, 90)), int(np.max(train_lengths))\n    )\n    logger.info(\"train_batches=%d\", len(train_loader))\n    logger.info(\"val_batches=%s\", len(val_loader) if val_loader is not None else 0)\n    logger.info(\"test_batches=%d\", len(test_loader))\n\n    log_section(\"TRAINING\")\n    model = RNA3DNet(len(BASES), D_MODEL, NHEAD, NUM_LAYERS, DROPOUT).to(DEVICE)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.SequentialLR(\n        optimizer,\n        schedulers=[\n            torch.optim.lr_scheduler.LinearLR(\n                optimizer,\n                start_factor=max(MIN_LR / max(LR, 1e-8), 0.2),\n                end_factor=1.0,\n                total_iters=max(WARMUP_EPOCHS, 1),\n            ),\n            torch.optim.lr_scheduler.CosineAnnealingLR(\n                optimizer,\n                T_max=max(1, MAX_EPOCHS - max(WARMUP_EPOCHS, 1)),\n                eta_min=MIN_LR,\n            ),\n        ],\n        milestones=[max(WARMUP_EPOCHS, 1)],\n    )\n\n    best_state = None\n    best_val = float(\"inf\")\n    bad_epochs = 0\n\n    for epoch in range(1, MAX_EPOCHS + 1):\n        CURRENT_EPOCH_FRACTION = (epoch - 1) / max(MAX_EPOCHS - 1, 1)\n        train_loss, train_time = run_epoch(train_loader, model, optimizer=optimizer)\n\n        if val_loader is not None:\n            val_loss, val_time = run_epoch(val_loader, model, optimizer=None)\n        else:\n            val_loss, val_time = train_loss, 0.0\n\n        scheduler.step()\n\n        logger.info(\n            \"EPOCH %02d | train_loss=%.5f | val_loss=%.5f | train_time=%.1fs | val_time=%.1fs | lr=%.2e | refine=%.2f\",\n            epoch, train_loss, val_loss, train_time, val_time, optimizer.param_groups[0][\"lr\"], CURRENT_EPOCH_FRACTION\n        )\n\n        if val_loss < best_val:\n            best_val = val_loss\n            bad_epochs = 0\n            best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n            logger.info(\"Novo melhor modelo salvo. best_val=%.5f\", best_val)\n        else:\n            bad_epochs += 1\n            logger.info(\"Sem melhora. patience=%d/%d\", bad_epochs, PATIENCE)\n            if bad_epochs >= PATIENCE:\n                logger.info(\"Early stopping acionado.\")\n                break\n\n    if best_state is not None:\n        model.load_state_dict(best_state)\n        logger.info(\"Melhor estado restaurado.\")\n\n    log_section(\"BUILD SUBMISSION\")\n    rows = []\n\n    for i, batch in enumerate(test_loader, start=1):\n        meta = batch[\"meta\"][0]\n        tid = meta[\"target_id\"]\n        logger.info(\"Inferência %d/%d -> target_id=%s length=%d\", i, len(test_loader), tid, meta[\"length\"])\n\n        preds5, ranked_candidates = sample_predictions_ranked(model, batch, n_samples=MC_SAMPLES, out_k=5)\n        for ridx, cand in enumerate(ranked_candidates[:5], start=1):\n            logger.info(\"target=%s rank=%d score=%.4f consensus=%.4f smooth=%.4f step_std=%.4f clash=%.4f compact=%.4f\", tid, ridx, cand[\"score\"], cand[\"features\"][\"consensus\"], cand[\"features\"][\"smooth\"], cand[\"features\"][\"step_std\"], cand[\"features\"][\"clash\"], cand[\"features\"][\"compact\"])\n        seq = meta[\"sequence\"]\n        resids = meta[\"resids\"]\n        L = meta[\"length\"]\n\n        for j in range(L):\n            row = {\n                \"ID\": f\"{tid}_{int(resids[j])}\",\n                \"resname\": seq[j],\n                \"resid\": int(resids[j]),\n            }\n            for k in range(5):\n                row[f\"x_{k+1}\"] = float(np.round(preds5[k][j, 0], 3))\n                row[f\"y_{k+1}\"] = float(np.round(preds5[k][j, 1], 3))\n                row[f\"z_{k+1}\"] = float(np.round(preds5[k][j, 2], 3))\n            rows.append(row)\n\n    submission = pd.DataFrame(rows)\n    ordered_cols = [\"ID\", \"resname\", \"resid\"]\n    for k in range(1, 6):\n        ordered_cols += [f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]\n\n    submission = submission[ordered_cols]\n    submission.to_csv(\"submission.csv\", index=False)\n\n    run_summary = {\n        \"version_tag\": VERSION_TAG,\n        \"device\": DEVICE,\n        \"train_samples\": len(train_samples),\n        \"val_samples\": len(val_samples) if val_loader is not None else 0,\n        \"test_samples\": len(test_samples),\n        \"best_val\": float(best_val),\n        \"batch_size\": BATCH_SIZE,\n        \"mc_samples\": MC_SAMPLES,\n        \"train_max_len\": TRAIN_MAX_LEN,\n        \"infer_chunk_len\": INFER_CHUNK_LEN,\n        \"infer_chunk_overlap\": INFER_CHUNK_OVERLAP,\n        \"cache_dir\": CACHE_DIR,\n    }\n    with open(\"run_summary.json\", \"w\", encoding=\"utf-8\") as f:\n        json.dump(run_summary, f, ensure_ascii=False, indent=2)\n\n    logger.info(\"submission.csv salvo com shape=%s\", submission.shape)\n    logger.info(\"Primeiras linhas:\\n%s\", submission.head())\n    logger.info(\"run_summary.json salvo\")\n    log_section(\"PIPELINE FINALIZADO COM SUCESSO\")\n\nexcept Exception as e:\n    log_section(\"PIPELINE ENCERRADO COM ERRO\")\n    logger.error(\"Erro fatal: %s\", e)\n    logger.error(\"Traceback completo:\\n%s\", traceback.format_exc())\n    raise\n","metadata":{},"outputs":[],"execution_count":null}]}