{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210,"isSourceIdPinned":false}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"04bb35a4","cell_type":"code","source":"# ============================================================\n# Stanford RNA 3D Folding Part 2 - V12.4 autoral focado em validação, generalização e submissão robusta\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 validado contra o sample oficial\n# - positional encoding fixa em 8192\n# - scheduler conservador orientado por validação/MIDVAL\n# - early stopping corrigido para preservar melhora via MIDVAL\n# - sliding windows de treino para sequências longas\n# - tuning para generalização e menos degradação intra-época\n# ============================================================\n\nimport os\nimport json\nimport hashlib\nimport math\nimport time\nimport random\nimport logging\nimport warnings\nimport traceback\nimport pickle\nimport re\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_v12_4\"\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 = \"V12_4_AUTHORAL_VALIDATION_RUNTIMEFIX\"\nMAX_EPOCHS = 4\nBATCH_SIZE = 2\nLR = 8.0e-5\nMIN_LR = 5.0e-6\nWEIGHT_DECAY = 3.2e-2\nD_MODEL = 160\nNHEAD = 4\nNUM_LAYERS = 3\nDROPOUT = 0.28\nMC_SAMPLES = 13\nGRAD_CLIP = 0.35\nPATIENCE = 2\nPE_BASE_MAX_LEN = 8192\nTRAIN_MAX_LEN = 1408\nTRAIN_MID_MAX_LEN = 2048\nINFER_CHUNK_LEN = 960\nINFER_CHUNK_OVERLAP = 288\nMAX_TRAIN_SAMPLES = 5716\nCACHE_DIR = \"/kaggle/working/rna3d_cache_v12_4\"\nUSE_SAMPLE_CACHE = True\nDIST_LOSS_MAX_POINTS = 112\nVAL_MAX_SAMPLES = 128\nPAIR_SEARCH_MAX_SPAN = 256\nPAIR_SEARCH_TOPK = 8\nLENGTH_BINS_FOR_CAP = 10\nPAIR_LOSS_W = 0.035\nSMOOTH_LOSS_W = 0.20\nSTEP_LOSS_W = 0.20\nSTERIC_LOSS_W_BASE = 0.04\nSTERIC_LOSS_W_FINAL = 0.085\nCURRENT_EPOCH_FRACTION = 0.0\nTRAIN_LENGTH_HARD_CAP = 6000\nSLIDING_WINDOW_MIN_LEN = 2048\nSLIDING_WINDOW_STRIDE_FRAC = 0.55\nSLIDING_WINDOW_MAX_PER_SAMPLE = 6\nSLIDING_WINDOW_JITTER = 96\nWARMUP_EPOCHS = 2\nPLATEAU_FACTOR = 0.50\nPLATEAU_PATIENCE = 1\nPLATEAU_THRESHOLD = 0.001\nINTRA_VAL_EVERY = 650\nINTRA_VAL_STEPS = 6\nBEST_DELTA = 1.0e-4\nPAIR_LOSS_W_START = 0.005\n\nLONG_SEQ_THRESHOLD = 1400\nCHUNK_MIN_CONTEXT = 96\nMAX_MC_SAMPLES_LONG = 8\nMAX_MC_SAMPLES_XLONG = 6\nRERANK_DIVERSITY_W = 0.17\nRERANK_VARIANCE_W = 0.05\nINFER_EDGE_WRITE_MIN = 96\nINFER_EDGE_WRITE_MAX = 224\nINFER_EDGE_WRITE_FRAC = 0.18\nRERANK_BOND_W = 0.10\nRERANK_CENTERLINE_W = 0.08\nRERANK_GLOBAL_SPREAD_W = 0.05\nLONG_SMOOTH_BLEND_W = 0.22\n\nPAIR_REFINEMENT_STEPS = 30\nPAIR_REFINEMENT_STEPS_LONG = 18\nPAIR_REFINEMENT_LR = 0.032\nPAIR_REFINEMENT_NOISE = 0.010\nPAIR_TARGET_DIST = 8.80\nBACKBONE_TARGET_STEP = 5.60\nOUTPUT_COORD_SCALE = 24.0\nMIN_GLOBAL_RG = 7.5\nPAIR_REFINEMENT_W = 0.26\nPAIR_RATIO_W = 0.22\nRERANK_PAIR_W = 0.24\nRERANK_SPREADMID_W = 0.06\nRERANK_MODE_BONUS = 0.06\nMMR_SCORE_W = 0.72\nMMR_DIVERSITY_W = 0.28\nCONSENSUS_CANDIDATE_W = 0.60\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nDEVICE_OBJ = torch.device(DEVICE)\nIS_CPU = DEVICE_OBJ.type == \"cpu\"\nIS_CUDA = DEVICE_OBJ.type == \"cuda\"\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 debug_list_input_roots(root=\"/kaggle/input\", max_dirs=40):\n    found = []\n    try:\n        for r, dirs, files in os.walk(root):\n            found.append((r, len(files)))\n            if len(found) >= max_dirs:\n                break\n    except Exception as e:\n        logger.warning(\"Falha ao listar %s: %s\", root, e)\n        return []\n    return found\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 notebook autoral V12.4: runtime fix do DEVICE, MIDVAL preservado, PE 8k fixo, scheduler conservador e sliding windows longos\")\n    logger.info(line)\n\n\ndef log_config_summary():\n    logger.info(\n        \"CONFIG SUMMARY | device=%s lr=%.2e min_lr=%.2e warmup=%d plateau_factor=%.2f dropout=%.2f wd=%.3f pair_w_end=%.3f steric=(%.3f->%.3f) intra_val_every=%d intra_val_steps=%d\",\n        DEVICE, LR, MIN_LR, WARMUP_EPOCHS, PLATEAU_FACTOR, DROPOUT, WEIGHT_DECAY,\n        PAIR_LOSS_W, STERIC_LOSS_W_BASE, STERIC_LOSS_W_FINAL, INTRA_VAL_EVERY, INTRA_VAL_STEPS,\n    )\n\n\ndef assert_runtime_config(train_samples_len: int):\n    if IS_CUDA and train_samples_len < MAX_TRAIN_SAMPLES:\n        logger.warning(\"CUDA active with reduced train_samples=%d; expected near full set=%d\", train_samples_len, MAX_TRAIN_SAMPLES)\n    if \"V12_4\" not in VERSION_TAG:\n        raise RuntimeError(f\"Unexpected VERSION_TAG at runtime: {VERSION_TAG}\")\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    keep_cols = [id_col, seq_col]\n    aux_map = {}\n    for src, dst in [(\"description\", \"description\"), (\"stoichiometry\", \"stoichiometry\"), (\"all_sequences\", \"all_sequences\"), (\"ligand_ids\", \"ligand_ids\"), (\"ligand_SMILES\", \"ligand_SMILES\")]:\n        col = find_col(df, [src])\n        if col is not None and col not in keep_cols:\n            keep_cols.append(col)\n            aux_map[col] = dst\n    out = df[keep_cols].copy()\n    out.columns = [\"target_id\", \"sequence\"] + [aux_map[c] for c in keep_cols[2:]]\n    out[\"target_id\"] = out[\"target_id\"].astype(str)\n    out[\"sequence\"] = out[\"sequence\"].astype(str).str.upper().str.replace(\" \", \"\", regex=False)\n    for col in out.columns:\n        if col not in [\"target_id\", \"sequence\"]:\n            out[col] = out[col].fillna(\"\").astype(str)\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    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_plain_col = find_col(df, [\"x\"])\n    y_plain_col = find_col(df, [\"y\"])\n    z_plain_col = find_col(df, [\"z\"])\n\n    # -------- Caso 1: 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 schema candidates -> ID=%s resname=%s resid=%s x=%s y=%s z=%s | wide target=%s seq=%s x=%d y=%d z=%d\",\n        df_name, id_col, resname_col, resid_col, x_plain_col, y_plain_col, z_plain_col, target_col, seq_col, len(x_cols), len(y_cols), len(z_cols)\n    )\n\n    has_plain_residue_cols = None not in [id_col, resname_col, resid_col, x_plain_col, y_plain_col, z_plain_col]\n    use_wide_schema = (\n        target_col is not None and\n        len(x_cols) > 1 and\n        len(x_cols) == len(y_cols) == len(z_cols)\n    )\n\n    if use_wide_schema:\n        rows = []\n        for _, row in df.iterrows():\n            raw_id = str(row[target_col])\n            tid = raw_id\n            if seq_col is not None and len(x_cols) > 1 and \"_\" in raw_id:\n                tid = raw_id.rsplit(\"_\", 1)[0]\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    # -------- Caso 2: formato longo --------\n    if has_plain_residue_cols:\n        out = df[[id_col, resname_col, resid_col, x_plain_col, y_plain_col, z_plain_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    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 infer_chain_break_flags(sequence, all_sequences_text=\"\"):\n    seq = str(sequence).upper()\n    L = len(seq)\n    flags = np.zeros(L, dtype=np.float32)\n    txt = str(all_sequences_text or \"\").upper()\n    if not txt:\n        return flags\n    pieces = [p for p in re.findall(r\"[ACGUTN]+\", txt) if p]\n    if len(pieces) <= 1:\n        return flags\n    joined = \"\".join(pieces)\n    if joined != seq:\n        return flags\n    pos = 0\n    for part in pieces[:-1]:\n        pos += len(part)\n        if 0 < pos < L:\n            flags[pos - 1] = 1.0\n    return flags\n\ndef build_context_vector(row, sequence):\n    seq = str(sequence).upper()\n    desc = str(row.get(\"description\", \"\") or \"\").lower()\n    sto = str(row.get(\"stoichiometry\", \"\") or \"\")\n    all_seq = str(row.get(\"all_sequences\", \"\") or \"\")\n    lig_ids = str(row.get(\"ligand_ids\", \"\") or \"\")\n    lig_smiles = str(row.get(\"ligand_SMILES\", \"\") or \"\")\n    chain_tokens = [p for p in re.findall(r\"[ACGUTN]+\", all_seq.upper()) if p]\n    chain_count = len(chain_tokens) if len(chain_tokens) > 1 else max(1, all_seq.count(\",\") + all_seq.count(\"|\") + all_seq.count(\";\") + 1 if all_seq else 1)\n    ligand_count = max(len([x for x in re.split(r\"[;,| ]+\", lig_ids) if x]), len([x for x in re.split(r\"[;,| ]+\", lig_smiles) if x]))\n    has_ligand = 1.0 if ligand_count > 0 or lig_smiles.strip() else 0.0\n    protein_hint = 1.0 if any(tok in desc for tok in [\"protein\", \"ribosome\", \"enzyme\", \"complex\", \"bound\", \"subunit\"]) else 0.0\n    state_hint = 1.0 if any(tok in desc for tok in [\"state\", \"conformation\", \"alternate\", \"apo\", \"holo\", \"active\", \"inactive\"]) else 0.0\n    multichain = 1.0 if chain_count > 1 else 0.0\n    long_seq = 1.0 if len(seq) >= 1200 else 0.0\n    return np.asarray([\n        min(1.0, math.log1p(len(seq)) / 9.0),\n        min(1.0, (chain_count - 1) / 4.0),\n        min(1.0, ligand_count / 4.0),\n        has_ligand,\n        protein_hint,\n        state_hint,\n        multichain,\n        long_seq,\n    ], dtype=np.float32)\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    row_map = {}\n    for row in seq_df.to_dict(\"records\"):\n        row_map[str(row[\"target_id\"])] = row\n    samples = []\n\n    if lbl_df is None:\n        for tid, seq in seq_map.items():\n            row = row_map.get(str(tid), {})\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                \"chain_breaks\": infer_chain_break_flags(seq, row.get(\"all_sequences\", \"\")),\n                \"context_vec\": build_context_vector(row, seq),\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        row = row_map.get(str(tid), {})\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            \"chain_breaks\": infer_chain_break_flags(seq[:L], row.get(\"all_sequences\", \"\"))[:L],\n            \"context_vec\": build_context_vector(row, seq[:L]),\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\n    hint_start = sample.get(\"crop_start_hint\")\n    if hint_start is not None:\n        start_lo = max(0, int(hint_start) - SLIDING_WINDOW_JITTER)\n        start_hi = min(L - crop_len, int(hint_start) + SLIDING_WINDOW_JITTER)\n        if start_hi < start_lo:\n            start_hi = start_lo\n        start = random.randint(start_lo, start_hi) if start_hi > start_lo else start_lo\n    else:\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(\"chain_breaks\") is not None:\n        out[\"chain_breaks\"] = sample[\"chain_breaks\"][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\ndef expand_training_samples_with_sliding_windows(samples, max_len=TRAIN_MAX_LEN, mid_len=TRAIN_MID_MAX_LEN):\n    expanded = []\n    augmented = 0\n    extra_views = 0\n    for sample in samples:\n        expanded.append(sample)\n        if sample.get(\"coords\") is None:\n            continue\n        L = len(sample[\"sequence\"])\n        if L < SLIDING_WINDOW_MIN_LEN:\n            continue\n        window_len = max_len if L > 2048 else mid_len\n        window_len = min(window_len, L)\n        if window_len >= L:\n            continue\n\n        stride = max(window_len // 3, int(window_len * SLIDING_WINDOW_STRIDE_FRAC))\n        stride = min(stride, max(window_len - 128, 1))\n        starts = list(range(0, max(L - window_len + 1, 1), max(1, stride)))\n        final_start = max(0, L - window_len)\n        if not starts or starts[-1] != final_start:\n            starts.append(final_start)\n        if len(starts) > SLIDING_WINDOW_MAX_PER_SAMPLE:\n            idxs = np.linspace(0, len(starts) - 1, num=SLIDING_WINDOW_MAX_PER_SAMPLE, dtype=int)\n            starts = [starts[i] for i in idxs.tolist()]\n\n        if len(starts) <= 1:\n            continue\n\n        augmented += 1\n        for wi, start in enumerate(starts):\n            if wi == 0:\n                continue\n            view = dict(sample)\n            view[\"crop_start_hint\"] = int(start)\n            view[\"crop_end_hint\"] = int(min(L, start + window_len))\n            view[\"window_index\"] = wi\n            view[\"window_parent_len\"] = L\n            expanded.append(view)\n            extra_views += 1\n\n    logger.info(\"Sliding-window augmentation -> base=%d expanded=%d long_augmented=%d extra_views=%d\", len(samples), len(expanded), augmented, extra_views)\n    return expanded\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\"{VERSION_TAG}_{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            if not isinstance(samples, list):\n                raise ValueError(\"cache invalido: objeto nao e lista\")\n            if len(samples) == 0 and lbl_df is not None:\n                raise ValueError(\"cache invalido: lista vazia para samples com labels\")\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        chain_breaks = np.asarray(s.get(\"chain_breaks\", np.zeros(L, dtype=np.float32)), dtype=np.float32)\n        if len(chain_breaks) != L:\n            chain_breaks = np.zeros(L, dtype=np.float32)\n        context_vec = np.asarray(s.get(\"context_vec\", np.zeros(8, dtype=np.float32)), dtype=np.float32)\n        if context_vec.shape != (8,):\n            context_vec = np.zeros(8, dtype=np.float32)\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            \"chain_breaks\": chain_breaks,\n            \"context_vec\": context_vec,\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    chain_breaks = 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    context_vec = torch.zeros((B, 8), dtype=torch.float32)\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        chain_breaks[i, :L] = torch.from_numpy(x[\"chain_breaks\"])\n        pos_norm[i, :L] = torch.from_numpy(x[\"pos_norm\"])\n        mask[i, :L] = True\n        context_vec[i] = torch.from_numpy(x[\"context_vec\"])\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            \"chain_breaks\": x[\"chain_breaks\"],\n            \"context_vec\": x[\"context_vec\"],\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        \"chain_breaks\": chain_breaks,\n        \"pos_norm\": pos_norm,\n        \"mask\": mask,\n        \"context_vec\": context_vec,\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.chain_break_mlp = nn.Sequential(\n            nn.Linear(1, d_model // 4),\n            nn.GELU(),\n            nn.Linear(d_model // 4, d_model),\n        )\n        self.context_mlp = nn.Sequential(\n            nn.Linear(8, d_model),\n            nn.GELU(),\n            nn.Linear(d_model, 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, chain_breaks, context_vec, 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        x = x + 0.16 * self.chain_break_mlp(chain_breaks.unsqueeze(-1))\n        x = x + 0.12 * self.context_mlp(context_vec).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        raw_coords = self.coord_head(x)\n        coords = torch.tanh(raw_coords / 3.2) * OUTPUT_COORD_SCALE\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, chain_breaks=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    if chain_breaks is not None:\n        edge_keep = (chain_breaks[:, :-1] < 0.5)\n        m2 = m2 & edge_keep\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_target = paired_flag * 0.90 + 0.05\n    pair_loss = nn.functional.binary_cross_entropy_with_logits(pair_logits, pair_target, 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    pair_w = PAIR_LOSS_W_START + (PAIR_LOSS_W - PAIR_LOSS_W_START) * float(CURRENT_EPOCH_FRACTION)\n    total = dist_loss + SMOOTH_LOSS_W * smooth + pair_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        \"pair_w\": float(pair_w),\n    }\n    return total, metrics\n\n\n# ============================================================\n# 11) LOOP DE TREINO\n# ============================================================\ndef run_epoch(loader, model, optimizer=None, max_steps=None, log_prefix=None):\n    training = optimizer is not None\n    model.train(training)\n    prefix = log_prefix or (\"TRAIN\" if training else \"VALID\")\n    total_loss = 0.0\n    total_n = 0\n    t0 = time.time()\n\n    for step, batch in enumerate(loader, start=1):\n        if max_steps is not None and step > max_steps:\n            break\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        chain_breaks = batch[\"chain_breaks\"].to(DEVICE)\n        context_vec = batch[\"context_vec\"].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, chain_breaks, context_vec, pos_norm, mask)\n            try:\n                loss, loss_parts = loss_fn(outputs, coords, mask, paired_flag, valid_mask=valid_mask, chain_breaks=chain_breaks)\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                prefix,\n                step, min(len(loader), max_steps or 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.024\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\n\ndef get_batch_pair_priors(batch):\n    mask = batch[\"mask\"][0].detach().cpu().numpy().astype(bool)\n    pair_map = batch[\"pair_map\"][0].detach().cpu().numpy()\n    pair_prob = batch[\"pair_prob\"][0].detach().cpu().numpy()\n    paired_flag = batch[\"paired_flag\"][0].detach().cpu().numpy() if \"paired_flag\" in batch else np.zeros_like(pair_prob)\n    return pair_map[mask], pair_prob[mask], paired_flag[mask]\n\n\ndef backbone_step_stats(pred):\n    pred = np.asarray(pred, dtype=np.float32)\n    if len(pred) < 2:\n        return {\n            \"step_mean\": 0.0,\n            \"step_median\": 0.0,\n            \"tiny_frac\": 1.0,\n            \"rg\": 0.0,\n        }\n    d = np.linalg.norm(pred[1:] - pred[:-1], axis=1)\n    rg = float(np.sqrt(np.mean(np.sum((pred - pred.mean(axis=0, keepdims=True)) ** 2, axis=1)) + 1e-8))\n    return {\n        \"step_mean\": float(np.mean(d)),\n        \"step_median\": float(np.median(d)),\n        \"tiny_frac\": float(np.mean(d < max(0.30, BACKBONE_TARGET_STEP * 0.35))),\n        \"rg\": rg,\n    }\n\ndef target_radius_of_gyration(L, pair_prob=None):\n    pair_frac = float(np.mean(pair_prob > 0.10)) if pair_prob is not None and len(pair_prob) else 0.0\n    Lf = max(4.0, float(L))\n    base = max(MIN_GLOBAL_RG, 2.40 + 0.92 * math.sqrt(Lf) + 0.055 * (Lf ** 0.72))\n    return float(base * (1.0 + 0.16 * pair_frac))\n\ndef mode_target_step(mode=\"base\", pair_prob=None, L=None):\n    pair_frac = float(np.mean(pair_prob > 0.10)) if pair_prob is not None and len(pair_prob) else 0.0\n    long_bonus = 0.0\n    if L is not None:\n        if L >= 3200:\n            long_bonus = 0.28\n        elif L >= 1600:\n            long_bonus = 0.16\n        elif L >= 800:\n            long_bonus = 0.08\n    mode_mult = {\n        \"base\": 1.00,\n        \"geom\": 1.03,\n        \"template\": 1.07,\n        \"stretch\": 1.12,\n    }.get(mode, 1.0)\n    return float(BACKBONE_TARGET_STEP * mode_mult * (1.0 + 0.10 * pair_frac) + long_bonus)\n\ndef enforce_backbone_step_profile(pred, target_step=BACKBONE_TARGET_STEP, strength=0.52, iters=2, chain_breaks=None):\n    pred = np.asarray(pred, dtype=np.float32).copy()\n    if len(pred) < 3:\n        return pred\n    target_step = float(target_step)\n    cb = np.asarray(chain_breaks, dtype=np.float32) if chain_breaks is not None and len(chain_breaks) == len(pred) else None\n    for _ in range(int(iters)):\n        delta = pred[1:] - pred[:-1]\n        step = np.linalg.norm(delta, axis=1) + 1e-8\n        if cb is not None:\n            valid_edge = cb[:-1] < 0.5\n        else:\n            valid_edge = np.ones_like(step, dtype=bool)\n        direction = delta / step[:, None]\n        desired = np.clip(step, target_step * 0.78, target_step * 1.28)\n        desired = np.where(step < target_step * 0.92, 0.55 * step + 0.45 * target_step, desired)\n        desired = np.where(step > target_step * 1.22, 0.72 * step + 0.28 * target_step, desired)\n        corr = (desired - step)[:, None] * direction * float(strength)\n        corr[~valid_edge] = 0.0\n        pred[:-1] -= 0.5 * corr\n        pred[1:]  += 0.5 * corr\n    pred -= pred.mean(axis=0, keepdims=True)\n    return pred.astype(np.float32)\n\ndef refine_local_distance_geometry(pred, target_step=BACKBONE_TARGET_STEP, chain_breaks=None, strength=0.22, iters=2):\n    pred = np.asarray(pred, dtype=np.float32).copy()\n    L = len(pred)\n    if L < 5:\n        return pred\n    cb = np.asarray(chain_breaks, dtype=np.float32) if chain_breaks is not None and len(chain_breaks) == L else np.zeros(L, dtype=np.float32)\n    target_step = float(target_step)\n    target2 = target_step * 1.78\n    target3 = target_step * 2.45\n    for _ in range(int(iters)):\n        for gap, target_d, w in [(1, target_step, 1.0), (2, target2, 0.55), (3, target3, 0.30)]:\n            for i in range(L - gap):\n                if cb[i:i + gap].max() >= 0.5:\n                    continue\n                vec = pred[i + gap] - pred[i]\n                dist = float(np.linalg.norm(vec)) + 1e-8\n                if not np.isfinite(dist):\n                    continue\n                direction = vec / dist if dist > 1e-8 else np.array([1.0, 0.0, 0.0], dtype=np.float32)\n                err = np.clip(target_d - dist, -0.75 * target_d, 0.75 * target_d)\n                push = strength * w * err\n                pred[i] -= 0.5 * push * direction\n                pred[i + gap] += 0.5 * push * direction\n        pred = expand_compact_regions(pred, min_step=target_step * 0.82)\n    pred -= pred.mean(axis=0, keepdims=True)\n    return pred.astype(np.float32)\n\ndef final_physical_rescale(pred, pair_prob=None, mode=\"base\", chain_breaks=None):\n    pred = np.asarray(pred, dtype=np.float32).copy()\n    L = len(pred)\n    if L < 2:\n        return pred\n    pred -= pred.mean(axis=0, keepdims=True)\n\n    target_step = mode_target_step(mode=mode, pair_prob=pair_prob, L=L)\n    d = np.linalg.norm(pred[1:] - pred[:-1], axis=1)\n    mean_step = float(np.mean(d)) if len(d) else target_step\n    med_step = float(np.median(d)) if len(d) else target_step\n    target_rg = target_radius_of_gyration(L, pair_prob=pair_prob) * ({\"stretch\":1.14, \"template\":1.08, \"geom\":1.03}.get(mode, 1.0))\n    rg = float(np.sqrt(np.mean(np.sum(pred ** 2, axis=1)) + 1e-8))\n\n    step_scale = target_step / max(0.65 * med_step + 0.35 * mean_step, 1e-6)\n    rg_scale = target_rg / max(rg, 1e-6)\n    tiny_frac = float(np.mean(d < max(0.35, target_step * 0.42))) if len(d) else 0.0\n    scale = 0.52 * step_scale + 0.48 * rg_scale\n    if mean_step < target_step * 0.80:\n        scale *= (1.10 + 0.16 * min(1.0, tiny_frac * 2.2))\n    if rg < target_rg * 0.78:\n        scale *= 1.08\n    scale = float(np.clip(scale, 0.92, 40.0))\n\n    pred *= scale\n    pred = expand_compact_regions(pred, min_step=target_step * (0.80 if mode == \"base\" else 0.88))\n    pred = refine_local_distance_geometry(pred, target_step=target_step, chain_breaks=chain_breaks, strength=0.20 if mode == \"base\" else 0.24, iters=2)\n    pred = enforce_backbone_step_profile(pred, target_step=target_step, strength=0.62 if mode in {\"template\", \"stretch\"} else 0.58, iters=4, chain_breaks=chain_breaks)\n    pred -= pred.mean(axis=0, keepdims=True)\n    return pred.astype(np.float32)\n\ndef calibrate_backbone_scale(pred, pair_prob=None, target_step=BACKBONE_TARGET_STEP, mode=\"base\"):\n    pred = np.asarray(pred, dtype=np.float32).copy()\n    L = len(pred)\n    if L < 2:\n        return pred\n\n    desired_step = mode_target_step(mode=mode, pair_prob=pair_prob, L=L)\n    d = np.linalg.norm(pred[1:] - pred[:-1], axis=1)\n    med = float(np.median(d)) if len(d) else desired_step\n    mean_step = float(np.mean(d)) if len(d) else desired_step\n    if med < 1e-6:\n        med = max(desired_step, 1.0)\n\n    scale_step = desired_step / med\n\n    centered = pred - pred.mean(axis=0, keepdims=True)\n    rg = float(np.sqrt(np.mean(np.sum(centered ** 2, axis=1)) + 1e-8))\n    desired_rg = target_radius_of_gyration(L, pair_prob=pair_prob) * ({\"stretch\":1.12, \"template\":1.06, \"geom\":1.02}.get(mode, 1.0))\n    scale_rg = desired_rg / max(rg, 1e-6)\n\n    tiny_frac = float(np.mean(d < max(0.30, desired_step * 0.38))) if len(d) else 0.0\n    blend = 0.50 + 0.28 * min(1.0, tiny_frac * 1.8)\n    scale = (1.0 - blend) * scale_step + blend * scale_rg\n    if mean_step < desired_step * 0.65:\n        scale *= 1.16\n    if rg < desired_rg * 0.70:\n        scale *= 1.10\n    scale = float(np.clip(scale, 0.92, 48.0))\n\n    pred = centered * scale\n    return pred.astype(np.float32)\n\ndef expand_compact_regions(pred, min_step=None):\n    pred = np.asarray(pred, dtype=np.float32).copy()\n    if len(pred) < 3:\n        return pred\n    min_step = float(min_step or (BACKBONE_TARGET_STEP * 0.72))\n    for _ in range(3):\n        delta = pred[1:] - pred[:-1]\n        step = np.linalg.norm(delta, axis=1) + 1e-8\n        bad = np.where(step < min_step)[0]\n        if len(bad) == 0:\n            break\n        for idx in bad:\n            direction = delta[idx] / step[idx]\n            if not np.all(np.isfinite(direction)) or np.linalg.norm(direction) < 1e-8:\n                if idx > 0:\n                    direction = pred[idx + 1] - pred[idx - 1]\n                    dn = np.linalg.norm(direction) + 1e-8\n                    direction = direction / dn\n                else:\n                    direction = np.array([1.0, 0.0, 0.0], dtype=np.float32)\n            push = 0.5 * (min_step - step[idx])\n            pred[idx] -= direction * push\n            pred[idx + 1] += direction * push\n    pred -= pred.mean(axis=0, keepdims=True)\n    return pred.astype(np.float32)\n\ndef pair_distance_ratio_score(pred, pair_map, pair_prob):\n    pred = np.asarray(pred, dtype=np.float32)\n    L = len(pred)\n    if L < 8:\n        return 1.0\n    pair_d = []\n    pair_w = []\n    far_d = []\n    for i in range(L):\n        j = int(pair_map[i]) if i < len(pair_map) else -1\n        if j > i and j < L:\n            dij = float(np.linalg.norm(pred[j] - pred[i]))\n            pair_d.append(dij)\n            pair_w.append(float(max(0.05, pair_prob[i] if i < len(pair_prob) else 0.5)))\n        k = min(L - 1, i + 12)\n        if k > i + 3:\n            far_d.append(float(np.linalg.norm(pred[k] - pred[i])))\n    if not pair_d or not far_d:\n        return 1.0\n    pair_mean = float(np.average(np.asarray(pair_d, dtype=np.float32), weights=np.asarray(pair_w, dtype=np.float32)))\n    far_med = float(np.median(np.asarray(far_d, dtype=np.float32))) + 1e-6\n    return pair_mean / far_med\n\ndef refine_geometry_with_pairs(pred, pair_map, pair_prob, steps=PAIR_REFINEMENT_STEPS, lr=PAIR_REFINEMENT_LR, strong=False, chain_breaks=None):\n    pred = np.asarray(pred, dtype=np.float32)\n    L = len(pred)\n    if L < 8:\n        return pred.copy()\n    pair_idx = [(i, int(pair_map[i]), float(pair_prob[i])) for i in range(min(L, len(pair_map))) if int(pair_map[i]) > i and int(pair_map[i]) < L]\n    if not pair_idx:\n        return pred.copy()\n\n    x = torch.tensor(pred, dtype=torch.float32, device=DEVICE)\n    x = x + torch.randn_like(x) * PAIR_REFINEMENT_NOISE\n    x = torch.nn.Parameter(x)\n    opt = torch.optim.Adam([x], lr=lr * (1.15 if strong else 1.0))\n\n    cb = np.asarray(chain_breaks, dtype=np.float32) if chain_breaks is not None and len(chain_breaks) == L else np.zeros(L, dtype=np.float32)\n    target_step = BACKBONE_TARGET_STEP\n    pair_target = PAIR_TARGET_DIST * (0.96 if strong else 1.0)\n    n_steps = steps if L < 2200 else max(8, PAIR_REFINEMENT_STEPS_LONG)\n\n    for _ in range(int(n_steps)):\n        opt.zero_grad(set_to_none=True)\n        dx = x[1:] - x[:-1]\n        step = torch.linalg.norm(dx, dim=-1) + 1e-6\n        edge_keep = torch.tensor((cb[:-1] < 0.5).astype(np.float32), dtype=torch.float32, device=x.device)\n        bond_loss = (((step - target_step).abs()) * edge_keep).sum() / (edge_keep.sum() + 1e-6)\n\n        sec = x[2:] - 2.0 * x[1:-1] + x[:-2]\n        edge2_keep = torch.tensor(((cb[:-2] < 0.5) & (cb[1:-1] < 0.5)).astype(np.float32), dtype=torch.float32, device=x.device)\n        smooth_raw = (sec.pow(2).sum(dim=-1) + 1e-6).sqrt() if len(sec) else x.new_tensor([])\n        smooth_loss = (smooth_raw * edge2_keep).sum() / (edge2_keep.sum() + 1e-6) if len(sec) else x.new_tensor(0.0)\n\n        tri = x[3:] - x[:-3]\n        tri_d = torch.linalg.norm(tri, dim=-1) + 1e-6 if len(x) > 3 else x.new_tensor([])\n        edge3_keep = torch.tensor(((cb[:-3] < 0.5) & (cb[1:-2] < 0.5) & (cb[2:-1] < 0.5)).astype(np.float32), dtype=torch.float32, device=x.device) if len(x) > 3 else x.new_tensor([])\n        tri_target = target_step * 2.45\n        tri_loss = (((tri_d - tri_target).abs()) * edge3_keep).sum() / (edge3_keep.sum() + 1e-6) if len(x) > 3 else x.new_tensor(0.0)\n\n        sel_i = torch.tensor([p[0] for p in pair_idx], dtype=torch.long, device=x.device)\n        sel_j = torch.tensor([p[1] for p in pair_idx], dtype=torch.long, device=x.device)\n        sel_w = torch.tensor([max(0.10, p[2]) for p in pair_idx], dtype=torch.float32, device=x.device)\n        pd = torch.linalg.norm(x[sel_i] - x[sel_j], dim=-1) + 1e-6\n        pair_loss = (((pd - pair_target).abs()) * sel_w).mean()\n\n        dmat = torch.cdist(x, x)\n        iu = torch.triu_indices(L, L, offset=3, device=x.device)\n        dd = dmat[iu[0], iu[1]]\n        clash_loss = torch.relu((0.78 if strong else 0.74) - dd).mean()\n\n        rg = torch.sqrt(torch.mean(torch.sum((x - x.mean(dim=0, keepdim=True)) ** 2, dim=-1)) + 1e-6)\n        rg_target = max(2.4, 0.09 * math.sqrt(L) * L ** 0.30)\n        spread_loss = (rg - rg_target).abs()\n\n        loss = 0.40 * bond_loss + 0.18 * smooth_loss + 0.10 * tri_loss + (PAIR_REFINEMENT_W * (1.20 if strong else 1.0)) * pair_loss + 0.16 * clash_loss + 0.07 * spread_loss\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_([x], 1.0)\n        opt.step()\n        with torch.no_grad():\n            x -= x.mean(dim=0, keepdim=True)\n\n    out = x.detach().cpu().numpy().astype(np.float32)\n    out = refine_local_distance_geometry(out, target_step=target_step, chain_breaks=cb, strength=0.18 if strong else 0.14, iters=2)\n    return out\n\n\ndef postprocess_candidate(pred, batch, sample_idx=0, total_samples=1, mode=\"base\"):\n    pair_map, pair_prob, paired_flag = get_batch_pair_priors(batch)\n    meta = batch.get(\"meta\", [{}])[0]\n    chain_breaks = np.asarray(meta.get(\"chain_breaks\", np.zeros(len(pred), dtype=np.float32)), dtype=np.float32)\n    context_vec = np.asarray(meta.get(\"context_vec\", np.zeros(8, dtype=np.float32)), dtype=np.float32)\n    is_multichain = bool(context_vec[6] > 0.5) if len(context_vec) > 6 else False\n    state_hint = float(context_vec[5]) if len(context_vec) > 5 else 0.0\n    pred = calibrate_backbone_scale(pred, pair_prob=pair_prob, mode=mode)\n    pred = final_physical_rescale(pred, pair_prob=pair_prob, mode=mode, chain_breaks=chain_breaks)\n    step_factor = 0.92 if is_multichain and mode in {\"template\", \"stretch\"} else (0.86 if mode in {\"template\", \"stretch\"} else 0.78)\n    pred = expand_compact_regions(pred, min_step=mode_target_step(mode=mode, pair_prob=pair_prob, L=len(pred)) * step_factor)\n    pred = smooth_coords_with_breaks(pred, chain_breaks=chain_breaks, window=max(5, adaptive_window(len(pred)) - (0 if mode == \"base\" else 2)))\n    if len(pred) >= 2600:\n        pred = 0.74 * pred + 0.26 * smooth_coords_with_breaks(pred, chain_breaks=chain_breaks, window=min(7, adaptive_window(len(pred))))\n    pred -= pred.mean(axis=0, keepdims=True)\n    pred = add_controlled_diversity(pred, sample_idx, total_samples)\n    if state_hint > 0.5 and mode in {\"stretch\", \"consensus_geom\"}:\n        pred = pred * 1.03\n\n    if mode in {\"geom\", \"template\", \"stretch\"}:\n        pred = refine_geometry_with_pairs(\n            pred,\n            pair_map=pair_map,\n            pair_prob=pair_prob,\n            steps=PAIR_REFINEMENT_STEPS if mode == \"geom\" else PAIR_REFINEMENT_STEPS + (8 if mode == \"template\" else 10),\n            lr=PAIR_REFINEMENT_LR * (1.00 if mode == \"geom\" else (0.92 if mode == \"template\" else 0.88)),\n            strong=(mode in {\"template\", \"stretch\"}),\n            chain_breaks=chain_breaks,\n        )\n        pred = calibrate_backbone_scale(pred, pair_prob=pair_prob, mode=mode)\n        pred = final_physical_rescale(pred, pair_prob=pair_prob, mode=mode, chain_breaks=chain_breaks)\n        pred = expand_compact_regions(pred, min_step=mode_target_step(mode=mode, pair_prob=pair_prob, L=len(pred)) * (0.90 if mode in {\"template\", \"stretch\"} else 0.82))\n        pred = smooth_coords_with_breaks(pred, chain_breaks=chain_breaks, window=min(7, max(3, adaptive_window(len(pred)) - 2)))\n        pred = refine_local_distance_geometry(pred, target_step=mode_target_step(mode=mode, pair_prob=pair_prob, L=len(pred)), chain_breaks=chain_breaks, strength=0.18 if mode in {\"template\", \"stretch\"} else 0.14, iters=1)\n        pred = enforce_backbone_step_profile(pred, target_step=mode_target_step(mode=mode, pair_prob=pair_prob, L=len(pred)), strength=0.62 if mode in {\"template\", \"stretch\"} else 0.58, iters=3, chain_breaks=chain_breaks)\n        pred -= pred.mean(axis=0, keepdims=True)\n    return pred.astype(np.float32)\n\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 smooth_coords_with_breaks(arr, chain_breaks=None, window=5):\n    if chain_breaks is None or len(arr) <= 2:\n        return smooth_coords(arr, window=window)\n    chain_breaks = np.asarray(chain_breaks, dtype=np.float32)\n    if len(chain_breaks) != len(arr):\n        return smooth_coords(arr, window=window)\n    out = np.asarray(arr, dtype=np.float32).copy()\n    start = 0\n    for i in range(len(arr) - 1):\n        if chain_breaks[i] >= 0.5:\n            seg_len = i + 1 - start\n            if seg_len >= 3:\n                out[start:i + 1] = smooth_coords(out[start:i + 1], window=min(window, max(3, seg_len)))\n            start = i + 1\n    seg_len = len(arr) - start\n    if seg_len >= 3:\n        out[start:] = smooth_coords(out[start:], window=min(window, max(3, seg_len)))\n    return out.astype(np.float32)\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).astype(np.float32)\n    center = np.median(stack, axis=0)\n    if len(stack) <= 2:\n        return center.astype(np.float32)\n    for _ in range(2):\n        d = np.mean(np.linalg.norm(stack - center[None, :, :], axis=2), axis=1)\n        scale = float(np.median(d)) + 1e-6\n        w = 1.0 / (1.0 + (d / scale) ** 2)\n        w = w / np.clip(w.sum(), 1e-6, None)\n        center = np.sum(stack * w[:, None, None], axis=0)\n    return center.astype(np.float32)\n\ndef build_consensus_candidate(raw_preds, batch, mode=\"template\"):\n    center = robust_center(raw_preds)\n    if len(raw_preds) >= 3:\n        d = np.asarray([consensus_score(p, center) for p in raw_preds], dtype=np.float32)\n        keep = max(3, int(np.ceil(len(raw_preds) * CONSENSUS_CANDIDATE_W)))\n        idx = np.argsort(d)[:keep]\n        center = np.mean(np.stack([raw_preds[i] for i in idx], axis=0), axis=0)\n    return postprocess_candidate(np.asarray(center, dtype=np.float32), batch, sample_idx=0, total_samples=max(1, len(raw_preds)), mode=mode)\n\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 bond_length_irregularity(pred):\n    if len(pred) < 2:\n        return 0.0\n    d = np.linalg.norm(pred[1:] - pred[:-1], axis=1)\n    target = float(np.median(d))\n    return float(np.mean(np.abs(d - target)))\n\ndef centerline_drift_score(pred):\n    if len(pred) < 7:\n        return 0.0\n    sm = smooth_coords(pred, window=min(adaptive_window(len(pred)), 7))\n    return float(np.mean(np.linalg.norm(pred - sm, axis=1)))\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\n\ndef target_step_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.mean(np.abs(d - BACKBONE_TARGET_STEP)))\n\ndef tiny_step_fraction(pred):\n    if len(pred) < 2:\n        return 1.0\n    d = np.linalg.norm(pred[1:] - pred[:-1], axis=1)\n    return float(np.mean(d < (BACKBONE_TARGET_STEP * 0.55)))\n\ndef rg_deviation_score(pred):\n    L = len(pred)\n    rg = np.sqrt(np.mean(np.sum((pred - pred.mean(axis=0, keepdims=True)) ** 2, axis=1)))\n    target = target_radius_of_gyration(L)\n    return float(abs(rg - target))\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\n\n\ndef rank_candidates(cands, batch=None):\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    bond_vals = [bond_length_irregularity(c[\"pred\"]) for c in cands]\n    centerline_vals = [centerline_drift_score(c[\"pred\"]) for c in cands]\n    target_step_vals = [target_step_score(c[\"pred\"]) for c in cands]\n    tiny_step_vals = [tiny_step_fraction(c[\"pred\"]) for c in cands]\n    rg_dev_vals = [rg_deviation_score(c[\"pred\"]) for c in cands]\n\n    if batch is not None:\n        pair_map, pair_prob, paired_flag = get_batch_pair_priors(batch)\n        pair_ratio_vals = [pair_distance_ratio_score(c[\"pred\"], pair_map, pair_prob) for c in cands]\n    else:\n        pair_ratio_vals = [1.0 for _ 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    bond_n = normalize_feature_list(bond_vals, reverse=True)\n    centerline_n = normalize_feature_list(centerline_vals, reverse=True)\n    pair_ratio_n = normalize_feature_list(pair_ratio_vals, reverse=True)\n    target_step_n = normalize_feature_list(target_step_vals, reverse=True)\n    tiny_step_n = normalize_feature_list(tiny_step_vals, reverse=True)\n    rg_dev_n = normalize_feature_list(rg_dev_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        mode = c.get(\"mode\", \"base\")\n        mode_bonus = RERANK_MODE_BONUS if mode in {\"geom\", \"template\", \"stretch\"} or str(mode).startswith(\"consensus_\") else 0.0\n        if mode == \"template\":\n            mode_bonus += 0.03\n        elif mode == \"stretch\":\n            mode_bonus += 0.02\n        elif str(mode).startswith(\"consensus_\"):\n            mode_bonus += 0.025\n        score = (\n            0.12 * consensus_n[i] +\n            0.09 * smooth_n[i] +\n            0.06 * step_n[i] +\n            0.18 * clash_n[i] +\n            0.05 * compact_n[i] +\n            0.18 * target_step_n[i] +\n            0.17 * tiny_step_n[i] +\n            0.15 * rg_dev_n[i] +\n            RERANK_BOND_W * bond_n[i] +\n            RERANK_CENTERLINE_W * centerline_n[i] +\n            RERANK_PAIR_W * pair_ratio_n[i] +\n            mode_bonus\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            \"bond_irregularity\": float(bond_vals[i]),\n            \"centerline_drift\": float(centerline_vals[i]),\n            \"pair_ratio\": float(pair_ratio_vals[i]),\n            \"target_step_err\": float(target_step_vals[i]),\n            \"tiny_step_frac\": float(tiny_step_vals[i]),\n            \"rg_dev\": float(rg_dev_vals[i]),\n            \"mode\": mode,\n        }\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(BACKBONE_TARGET_STEP * 1.15, np.sqrt(L) * min_div_frac)\n    score_vals = np.asarray([float(c.get(\"score\", 0.0)) for c in cands], dtype=np.float32)\n    score_n = normalize_feature_list(score_vals, reverse=False)\n\n    selected = []\n    remaining = list(range(len(cands)))\n    while remaining and len(selected) < k:\n        if not selected:\n            best_idx = max(remaining, key=lambda idx: float(score_n[idx]))\n            selected.append(cands[best_idx])\n            remaining.remove(best_idx)\n            continue\n\n        best_idx = None\n        best_obj = -1e18\n        best_dist = 0.0\n        for idx in remaining:\n            dmin = min(pairwise_candidate_distance(cands[idx][\"pred\"], s[\"pred\"]) for s in selected)\n            div_term = min(1.0, dmin / max(scale, 1e-6))\n            obj = MMR_SCORE_W * float(score_n[idx]) + MMR_DIVERSITY_W * float(div_term)\n            if dmin < scale * 0.55:\n                obj -= 0.35\n            if obj > best_obj:\n                best_obj = obj\n                best_idx = idx\n                best_dist = dmin\n        if best_idx is None:\n            break\n        if best_dist >= scale * 0.72 or len(selected) == 0 or len(remaining) <= (k - len(selected)):\n            selected.append(cands[best_idx])\n            remaining.remove(best_idx)\n        else:\n            selected.append(cands[best_idx])\n            remaining.remove(best_idx)\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        filler_idx = 0\n        while len(selected) < k:\n            base = cands[filler_idx % len(cands)]\n            cloned = {\n                \"pred\": add_controlled_diversity(base[\"pred\"], filler_idx + len(selected), k),\n                \"score\": float(base.get(\"score\", 0.0)),\n                \"features\": dict(base.get(\"features\", {}))\n            }\n            selected.append(cloned)\n            filler_idx += 1\n    return selected[:k]\n\ndef choose_mc_samples(seq_len, base=MC_SAMPLES):\n    if seq_len >= 5200:\n        return max(4, min(MAX_MC_SAMPLES_XLONG, base))\n    if seq_len >= 3600:\n        return max(5, min(MAX_MC_SAMPLES_XLONG + 1, base))\n    if seq_len >= 1800:\n        return max(7, min(MAX_MC_SAMPLES_LONG, base))\n    if seq_len <= 480:\n        return min(base + 2, 16)\n    if seq_len <= 960:\n        return min(base + 1, 15)\n    return int(base)\n\ndef blend_weights_1d(length):\n    if length <= 1:\n        return np.ones((length,), dtype=np.float32)\n    x = np.linspace(0.0, 1.0, length, dtype=np.float32)\n    w = 1.0 - np.abs(2.0 * x - 1.0)\n    w = 0.25 + 0.75 * w\n    return w.astype(np.float32)\n\ndef choose_infer_chunk_params(full_len, base_chunk=INFER_CHUNK_LEN, base_overlap=INFER_CHUNK_OVERLAP):\n    if full_len >= 12000:\n        return 640, 192\n    if full_len >= 8000:\n        return 704, 192\n    if full_len >= 4800:\n        return 768, 224\n    if full_len >= 2600:\n        return base_chunk, base_overlap\n    return min(1024, max(base_chunk, full_len)), min(base_overlap, max(96, full_len // 6))\n\ndef _slice_tensor_for_infer(key, value, start, end):\n    if not isinstance(value, torch.Tensor):\n        return value\n    if key in {\"pair_map\", \"pair_prob\", \"pair_type\"}:\n        if value.ndim == 3:\n            return value[:, start:end, start:end].contiguous()\n        if value.ndim == 2:\n            return value[:, start:end].contiguous()\n        if value.ndim == 1:\n            return value[start:end].contiguous()\n        return value.contiguous()\n    if key in {\"seq_ids\", \"pos_norm\", \"mask\", \"coords\", \"valid_mask\", \"paired_flag\", \"chain_breaks\"}:\n        if value.ndim >= 2:\n            return value[:, start:end].contiguous()\n        if value.ndim == 1:\n            return value[start:end].contiguous()\n        return value.contiguous()\n    return value\n\ndef slice_infer_batch(batch, start, end):\n    sub = {}\n    for key, value in batch.items():\n        if key == \"meta\":\n            meta = dict(value[0])\n            meta[\"length\"] = int(end - start)\n            meta[\"sequence\"] = meta[\"sequence\"][start:end]\n            meta[\"resids\"] = meta[\"resids\"][start:end]\n            if \"chain_breaks\" in meta:\n                meta[\"chain_breaks\"] = meta[\"chain_breaks\"][start:end]\n            sub[\"meta\"] = [meta]\n        elif isinstance(value, torch.Tensor):\n            sub[key] = _slice_tensor_for_infer(key, value, start, end)\n        else:\n            sub[key] = value\n    return sub\n\ndef infer_coords_single_pass(model, batch):\n\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    chain_breaks = batch[\"chain_breaks\"].to(DEVICE)\n    context_vec = batch[\"context_vec\"].to(DEVICE)\n    pos_norm = batch[\"pos_norm\"].to(DEVICE)\n    mask = batch[\"mask\"].to(DEVICE)\n    with torch.no_grad():\n        outputs = model(seq_ids, pair_map, pair_prob, pair_type, chain_breaks, context_vec, pos_norm, mask)\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    return pred\n\n\n\ndef infer_coords_chunked(model, batch, chunk_len=INFER_CHUNK_LEN, overlap=INFER_CHUNK_OVERLAP):\n    full_len = int(batch[\"mask\"][0].sum().item())\n    chunk_len, overlap = choose_infer_chunk_params(full_len, base_chunk=chunk_len, base_overlap=overlap)\n\n    if full_len <= chunk_len:\n        return infer_coords_single_pass(model, batch)\n\n    step = max(64, chunk_len - overlap)\n    acc = np.zeros((full_len, 3), dtype=np.float32)\n    den = np.zeros((full_len, 1), dtype=np.float32)\n\n    starts = list(range(0, full_len, step))\n    final_start = max(0, full_len - chunk_len)\n    if not starts or starts[-1] != final_start:\n        starts.append(final_start)\n\n    total_windows = len(starts)\n    for wi, start in enumerate(starts):\n        end = min(full_len, start + chunk_len)\n        start = max(0, end - chunk_len)\n\n        sub_batch = slice_infer_batch(batch, start, end)\n        pred = infer_coords_single_pass(model, sub_batch)\n        local_len = len(pred)\n        w = blend_weights_1d(local_len).reshape(-1, 1)\n\n        edge_margin = int(max(INFER_EDGE_WRITE_MIN, min(INFER_EDGE_WRITE_MAX, overlap * (INFER_EDGE_WRITE_FRAC + 0.55))))\n        left_guard = 0 if wi == 0 else min(edge_margin, max(0, local_len // 2 - 1))\n        right_guard = 0 if wi == total_windows - 1 else min(edge_margin, max(0, local_len // 2 - 1))\n\n        write_start_local = min(left_guard, max(0, local_len - 1))\n        write_end_local = max(write_start_local + 1, local_len - right_guard)\n\n        acc[start + write_start_local:start + write_end_local] += pred[write_start_local:write_end_local] * w[write_start_local:write_end_local]\n        den[start + write_start_local:start + write_end_local] += w[write_start_local:write_end_local]\n\n        if IS_CUDA and full_len >= 4000:\n            torch.cuda.empty_cache()\n\n    pred = acc / np.clip(den, 1e-6, None)\n    pred = pred.astype(np.float32)\n    if full_len >= 2600:\n        sm = smooth_coords(pred, window=min(adaptive_window(full_len), 7))\n        pred = ((1.0 - LONG_SMOOTH_BLEND_W) * pred + LONG_SMOOTH_BLEND_W * sm).astype(np.float32)\n    return pred.astype(np.float32)\n\ndef normalize_sample_submission_df(sample_sub_df):\n    if sample_sub_df is None or sample_sub_df.empty:\n        raise ValueError(\"sample_submission.csv ausente ou vazio\")\n    sample = sample_sub_df.copy()\n    required_cols = [\"ID\", \"resname\", \"resid\"]\n    for k in range(1, 6):\n        required_cols += [f\"x_{k}\", f\"y_{k}\", f\"z_{k}\"]\n    missing = [c for c in required_cols if c not in sample.columns]\n    if missing:\n        raise ValueError(f\"sample_submission.csv fora do schema oficial: faltando {missing}\")\n    sample = sample[required_cols].copy()\n    sample[\"ID\"] = sample[\"ID\"].astype(str)\n    sample[\"resname\"] = sample[\"resname\"].astype(str).str.upper()\n    sample[\"resid\"] = pd.to_numeric(sample[\"resid\"], errors=\"raise\").astype(int)\n    if sample[\"ID\"].duplicated().any():\n        dup_ids = sample.loc[sample[\"ID\"].duplicated(), \"ID\"].astype(str).head(5).tolist()\n        raise ValueError(f\"sample_submission.csv possui IDs duplicados: {dup_ids}\")\n    return sample\n\ndef align_submission_to_sample(submission, sample_sub_df):\n    sample = normalize_sample_submission_df(sample_sub_df)\n    sub = submission.copy()\n    required_cols = list(sample.columns)\n    coord_cols = [c for c in required_cols if c not in [\"ID\", \"resname\", \"resid\"]]\n    missing_cols = [c for c in [\"ID\"] + coord_cols if c not in sub.columns]\n    if missing_cols:\n        raise ValueError(f\"submission intermediária sem colunas obrigatórias: {missing_cols}\")\n    sub = sub[[\"ID\"] + coord_cols].copy()\n    sub[\"ID\"] = sub[\"ID\"].astype(str)\n    if sub[\"ID\"].duplicated().any():\n        dup_ids = sub.loc[sub[\"ID\"].duplicated(), \"ID\"].astype(str).head(5).tolist()\n        raise ValueError(f\"Predição gerou IDs duplicados: {dup_ids}\")\n    for col in coord_cols:\n        sub[col] = pd.to_numeric(sub[col], errors=\"coerce\")\n    merged = sample[[\"ID\", \"resname\", \"resid\"]].merge(sub, on=\"ID\", how=\"left\")\n    missing_ids = merged.loc[merged[coord_cols].isna().all(axis=1), \"ID\"].astype(str).tolist()\n    if missing_ids:\n        raise ValueError(f\"Faltam predições para {len(missing_ids)} IDs do sample_submission. Exemplos: {missing_ids[:5]}\")\n    for col in coord_cols:\n        if merged[col].isna().any():\n            bad_ids = merged.loc[merged[col].isna(), \"ID\"].astype(str).head(5).tolist()\n            raise ValueError(f\"Coluna {col} contém NaN após alinhamento. Exemplos de IDs: {bad_ids}\")\n        merged[col] = merged[col].astype(float)\n    return merged[required_cols]\n\ndef validate_submission_matches_sample(submission, sample_sub_df):\n    sample = normalize_sample_submission_df(sample_sub_df)\n    required_cols = list(sample.columns)\n    if list(submission.columns) != required_cols:\n        raise ValueError(f\"Colunas finais da submission não batem com o schema oficial: {list(submission.columns)}\")\n    if len(submission) != len(sample):\n        raise ValueError(f\"Quantidade de linhas incorreta: submission={len(submission)} sample={len(sample)}\")\n    if submission[\"ID\"].astype(str).tolist() != sample[\"ID\"].astype(str).tolist():\n        raise ValueError(\"A ordem dos IDs da submission não corresponde ao sample_submission oficial\")\n    if submission[[\"resname\", \"resid\"]].reset_index(drop=True).equals(sample[[\"resname\", \"resid\"]].reset_index(drop=True)) is False:\n        raise ValueError(\"As colunas resname/resid devem reproduzir exatamente o sample_submission oficial\")\n    coord_cols = [c for c in required_cols if c not in [\"ID\", \"resname\", \"resid\"]]\n    if submission[coord_cols].isna().any().any():\n        raise ValueError(\"submission final contém NaN nas coordenadas\")\n    return submission[required_cols].copy()\n\ndef sample_submission_target_map(sample_sub_df):\n    sample = normalize_sample_submission_df(sample_sub_df)\n    target_map = {}\n    for row in sample[[\"ID\", \"resname\", \"resid\"]].itertuples(index=False):\n        id_str = str(row.ID)\n        if \"_\" not in id_str:\n            raise ValueError(f\"ID inválido no sample_submission: {id_str}\")\n        tid, resid_txt = id_str.rsplit(\"_\", 1)\n        resid_from_id = int(resid_txt)\n        if resid_from_id != int(row.resid):\n            raise ValueError(f\"sample_submission inconsistente para {id_str}: resid coluna={row.resid} resid ID={resid_from_id}\")\n        target_map.setdefault(tid, []).append({\n            \"ID\": id_str,\n            \"resname\": str(row.resname).upper(),\n            \"resid\": int(row.resid),\n        })\n    return target_map\n\ndef empty_submission_from_sample(sample_sub_df):\n    submission = normalize_sample_submission_df(sample_sub_df).copy()\n    coord_cols = [c for c in submission.columns if c not in [\"ID\", \"resname\", \"resid\"]]\n    for col in coord_cols:\n        submission[col] = np.nan\n    return submission\n\n\n\n\n\n\ndef sample_predictions_ranked(model, batch, n_samples=MC_SAMPLES, out_k=5):\n    full_len = int(batch[\"mask\"][0].sum().item())\n    n_samples = choose_mc_samples(full_len, base=n_samples)\n    chunk_len, overlap = choose_infer_chunk_params(full_len)\n\n    model.eval()\n    model.apply(enable_mc_dropout)\n\n    cands = []\n    candidate_modes = [\"base\", \"geom\", \"template\", \"stretch\"]\n\n    for sidx in range(n_samples):\n        pred = infer_coords_chunked(model, batch, chunk_len=chunk_len, overlap=overlap)\n        pred = np.asarray(pred, dtype=np.float32)\n\n        for mode in candidate_modes:\n            proc = postprocess_candidate(pred.copy(), batch, sample_idx=sidx, total_samples=n_samples, mode=mode)\n            stats = backbone_step_stats(proc)\n            cands.append({\n                \"pred\": proc,\n                \"variance\": 0.0,\n                \"spread\": stats[\"rg\"],\n                \"mode\": mode,\n                \"step_mean\": stats[\"step_mean\"],\n                \"step_median\": stats[\"step_median\"],\n                \"tiny_frac\": stats[\"tiny_frac\"],\n            })\n\n    raw_preds = [c[\"pred\"] for c in cands]\n    raw_preds = diversify_if_collapsed(raw_preds, target_k=max(out_k, len(raw_preds)))\n    consensus_mode = \"template\" if full_len <= 3600 else \"geom\"\n    consensus_pred = build_consensus_candidate(raw_preds, batch, mode=consensus_mode)\n    cands.append({\n        \"pred\": consensus_pred,\n        \"variance\": 0.0,\n        \"spread\": backbone_step_stats(consensus_pred)[\"rg\"],\n        \"mode\": f\"consensus_{consensus_mode}\",\n        \"step_mean\": backbone_step_stats(consensus_pred)[\"step_mean\"],\n        \"step_median\": backbone_step_stats(consensus_pred)[\"step_median\"],\n        \"tiny_frac\": backbone_step_stats(consensus_pred)[\"tiny_frac\"],\n    })\n    raw_preds = [c[\"pred\"] for c in cands]\n    center = robust_center(raw_preds)\n    variance_vals = [float(np.mean(np.linalg.norm(p - center, axis=1))) for p in raw_preds]\n    spread_vals = [float(np.sqrt(np.mean(np.sum((p - p.mean(axis=0, keepdims=True)) ** 2, axis=1)))) for p in raw_preds]\n\n    for i, p in enumerate(raw_preds):\n        stats = backbone_step_stats(p)\n        cands[i][\"pred\"] = p\n        cands[i][\"variance\"] = variance_vals[i]\n        cands[i][\"spread\"] = spread_vals[i]\n        cands[i][\"step_mean\"] = stats[\"step_mean\"]\n        cands[i][\"step_median\"] = stats[\"step_median\"]\n        cands[i][\"tiny_frac\"] = stats[\"tiny_frac\"]\n\n    cands = rank_candidates(cands, batch=batch)\n\n    variance_norm = normalize_feature_list([c.get(\"variance\", 0.0) for c in cands], reverse=False)\n    spread_norm = normalize_feature_list([c.get(\"spread\", 0.0) for c in cands], reverse=False)\n    spread_mid_penalty = np.abs(spread_norm - np.median(spread_norm)) if len(spread_norm) else np.asarray([], dtype=np.float32)\n\n    for i, c in enumerate(cands):\n        score = c[\"score\"] + RERANK_DIVERSITY_W * variance_norm[i]\n        score -= RERANK_VARIANCE_W * max(0.0, variance_norm[i] - 0.84)\n        if len(spread_mid_penalty):\n            score -= RERANK_GLOBAL_SPREAD_W * float(spread_mid_penalty[i])\n        score -= 0.16 * max(0.0, c.get(\"tiny_frac\", 0.0) - 0.06)\n        score -= 0.09 * abs(c.get(\"step_mean\", 0.0) - mode_target_step(mode=c.get(\"mode\", \"base\"), L=full_len)) / max(BACKBONE_TARGET_STEP, 1e-6)\n        score += 0.08 * min(1.0, c.get(\"spread\", 0.0) / max(target_radius_of_gyration(full_len), 1e-6)) if \"full_len\" in locals() else 0.0\n        c[\"score\"] = float(score)\n        c.setdefault(\"features\", {})\n        c[\"features\"][\"variance\"] = float(c.get(\"variance\", 0.0))\n        c[\"features\"][\"spread\"] = float(c.get(\"spread\", 0.0))\n        c[\"features\"][\"chunk_len\"] = int(chunk_len)\n        c[\"features\"][\"step_mean\"] = float(c.get(\"step_mean\", 0.0))\n        c[\"features\"][\"step_median\"] = float(c.get(\"step_median\", 0.0))\n        c[\"features\"][\"tiny_frac\"] = float(c.get(\"tiny_frac\", 0.0))\n\n    cands.sort(key=lambda x: x[\"score\"], reverse=True)\n\n    min_div = 0.16 if full_len < 1500 else (0.13 if full_len < 3200 else 0.10)\n    best = select_diverse_top_k(cands, k=out_k, min_div_frac=min_div)\n    preds = [c[\"pred\"] for c in best]\n\n    while len(preds) < out_k:\n        base = cands[min(len(preds), len(cands) - 1)][\"pred\"].copy()\n        filler = add_controlled_diversity(base, len(preds), out_k)\n        filler = calibrate_backbone_scale(filler, pair_prob=get_batch_pair_priors(batch)[1], mode=\"stretch\")\n        filler = final_physical_rescale(filler, pair_prob=get_batch_pair_priors(batch)[1], mode=\"stretch\")\n        preds.append(filler)\n\n    return preds, cands\n\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    if any(p is None for p in [train_seq_path, train_lbl_path, test_seq_path, sample_sub_path]):\n        logger.warning(\"Arquivos obrigatorios nao encontrados. Primeiros diretorios visiveis em /kaggle/input:\")\n        for root_path, nfiles in debug_list_input_roots(\"/kaggle/input\", max_dirs=40):\n            logger.warning(\"INPUT DIR -> %s | files=%d\", root_path, nfiles)\n        logger.warning(\"Se estiver no Kaggle, confirme se a competicao foi anexada na aba Input.\")\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    sample_sub_df = normalize_sample_submission_df(sample_sub_df)\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 IS_CPU and 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    else:\n        logger.info(\"GPU/full mode active, keeping train_samples -> %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    train_samples = expand_training_samples_with_sliding_windows(train_samples)\n\n    if IS_CPU and 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=IS_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=IS_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=IS_CUDA,\n    )\n\n    if len(train_samples) == 0:\n        raise RuntimeError(\"train_samples ficou vazio apos o parsing. Verifique se train_labels foi interpretado no schema correto e se os target_id batem com train_sequences.\")\n    if val_seq_df is not None and val_lbl_df is not None and len(val_samples) == 0:\n        raise RuntimeError(\"val_samples ficou vazio apos o parsing. Verifique validation_labels, validation_sequences e o alinhamento de target_id/resid.\")\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    assert_runtime_config(len(train_samples))\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    warmup_scheduler = torch.optim.lr_scheduler.LinearLR(\n        optimizer,\n        start_factor=max(MIN_LR / max(LR, 1e-8), 0.25),\n        end_factor=1.0,\n        total_iters=max(WARMUP_EPOCHS, 1),\n    )\n    plateau_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer,\n        mode=\"min\",\n        factor=PLATEAU_FACTOR,\n        patience=PLATEAU_PATIENCE,\n        threshold=PLATEAU_THRESHOLD,\n        threshold_mode=\"rel\",\n        cooldown=0,\n        min_lr=MIN_LR,\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        intra_val_loss = None\n        intra_val_time = 0.0\n        improved_this_epoch = False\n        if val_loader is not None and len(val_loader) > 0:\n            intra_val_loss, intra_val_time = run_epoch(val_loader, model, optimizer=None, max_steps=INTRA_VAL_STEPS, log_prefix=\"MIDVAL\")\n            logger.info(\"MID-EPOCH %02d | mid_val_loss=%.5f | mid_val_time=%.1fs | steps=%d\", epoch, intra_val_loss, intra_val_time, INTRA_VAL_STEPS)\n            if intra_val_loss < best_val - BEST_DELTA:\n                best_val = intra_val_loss\n                bad_epochs = 0\n                improved_this_epoch = True\n                best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n                logger.info(\"Novo melhor modelo salvo por MIDVAL. best_val=%.5f\", best_val)\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        prev_lr = optimizer.param_groups[0][\"lr\"]\n        monitor_val = float(val_loss)\n        if intra_val_loss is not None:\n            # Be conservative with the tiny validation set: use the better of mid/full validation\n            monitor_val = min(float(val_loss), float(intra_val_loss))\n\n        if epoch <= max(WARMUP_EPOCHS, 1):\n            warmup_scheduler.step()\n        else:\n            plateau_scheduler.step(monitor_val)\n        new_lr = optimizer.param_groups[0][\"lr\"]\n\n        logger.info(\n            \"EPOCH %02d | train_loss=%.5f | val_loss=%.5f | train_time=%.1fs | val_time=%.1fs | lr=%.2e->%.2e | refine=%.2f | mid_val=%s | monitor_val=%.5f\",\n            epoch, train_loss, val_loss, train_time, val_time, prev_lr, new_lr, CURRENT_EPOCH_FRACTION,\n            f\"{intra_val_loss:.5f}\" if intra_val_loss is not None else \"NA\", monitor_val,\n        )\n\n        if monitor_val < best_val - BEST_DELTA:\n            best_val = monitor_val\n            bad_epochs = 0\n            improved_this_epoch = True\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 monitor_val=%.5f\", best_val, monitor_val)\n        else:\n            if improved_this_epoch:\n                logger.info(\"Epoch %02d manteve o melhor estado via MIDVAL; patience preservada.\", epoch)\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    sample_target_map = sample_submission_target_map(sample_sub_df)\n    submission = empty_submission_from_sample(sample_sub_df)\n    coord_cols = [c for c in submission.columns if c not in [\"ID\", \"resname\", \"resid\"]]\n    seen_targets = set()\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        if tid in seen_targets:\n            raise ValueError(f\"target_id duplicado no test_loader: {tid}\")\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 mode=%s consensus=%.4f smooth=%.4f step_std=%.4f step_mean=%.4f tiny=%.4f clash=%.4f compact=%.4f variance=%.4f\", tid, ridx, cand[\"score\"], cand[\"features\"].get(\"mode\",\"base\"), cand[\"features\"][\"consensus\"], cand[\"features\"][\"smooth\"], cand[\"features\"][\"step_std\"], cand[\"features\"].get(\"step_mean\",0.0), cand[\"features\"].get(\"tiny_frac\",0.0), cand[\"features\"][\"clash\"], cand[\"features\"][\"compact\"], cand[\"features\"].get(\"variance\", 0.0))\n        template_rows = sample_target_map.get(tid)\n        if template_rows is None:\n            raise ValueError(f\"target_id {tid} não encontrado no sample_submission oficial\")\n        L = meta[\"length\"]\n        if len(template_rows) != L:\n            raise ValueError(f\"Comprimento inconsistente para {tid}: modelo={L} sample_submission={len(template_rows)}\")\n\n        ids = [tpl[\"ID\"] for tpl in template_rows]\n        mask = submission[\"ID\"].astype(str).isin(ids)\n        if int(mask.sum()) != L:\n            raise ValueError(f\"Falha ao localizar linhas do sample_submission para {tid}: esperadas={L} encontradas={int(mask.sum())}\")\n        fill_rows = []\n        for j, tpl in enumerate(template_rows):\n            row = {\"ID\": tpl[\"ID\"]}\n            for k in range(5):\n                row[f\"x_{k+1}\"] = float(preds5[k][j, 0])\n                row[f\"y_{k+1}\"] = float(preds5[k][j, 1])\n                row[f\"z_{k+1}\"] = float(preds5[k][j, 2])\n            fill_rows.append(row)\n        fill_df = pd.DataFrame(fill_rows)\n        submission.loc[mask, coord_cols] = fill_df.set_index(\"ID\").loc[submission.loc[mask, \"ID\"].astype(str), coord_cols].to_numpy()\n        seen_targets.add(tid)\n\n    missing_targets = sorted(set(sample_target_map.keys()) - seen_targets)\n    if missing_targets:\n        raise ValueError(f\"Faltaram targets do sample_submission na inferência: {missing_targets[:5]}\")\n    submission = validate_submission_matches_sample(submission, sample_sub_df)\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_base\": MC_SAMPLES,\n        \"long_seq_threshold\": LONG_SEQ_THRESHOLD,\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    if sample_sub_df is not None:\n        logger.info(\"sample_submission shape=%s | ids_sample=%d | ids_submission=%d\", sample_sub_df.shape, sample_sub_df[\"ID\"].astype(str).nunique(), submission[\"ID\"].astype(str).nunique())\n        logger.info(\"missing_ids_after_align=%d\", int((~sample_sub_df[\"ID\"].astype(str).isin(submission[\"ID\"].astype(str))).sum()))\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}]}