{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"colab":{"name":"test LLM","provenance":[]}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport re\nfrom collections import Counter\nimport unicodedata\nfrom sklearn.model_selection import train_test_split\nimport json\nfrom tqdm.notebook import tqdm\nfrom huggingface_hub import HfApi, login, hf_hub_download\nimport emoji\nimport ast\nimport random\nimport warnings\nimport os\nimport time\nfrom functools import partial\nfrom datasets import load_dataset\nimport nltk\nimport re\nfrom transformers import AutoModelForMaskedLM, AutoTokenizer\nimport math\n\nfrom torch.utils.checkpoint import checkpoint\nimport torch\nfrom torch import optim\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, Subset, DataLoader\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:38:47.886323Z","iopub.execute_input":"2026-02-03T13:38:47.886748Z","iopub.status.idle":"2026-02-03T13:38:47.893942Z","shell.execute_reply.started":"2026-02-03T13:38:47.886714Z","shell.execute_reply":"2026-02-03T13:38:47.893074Z"},"id":"01esU1q1p2Vw"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:38:47.895472Z","iopub.execute_input":"2026-02-03T13:38:47.89578Z","iopub.status.idle":"2026-02-03T13:38:47.908348Z","shell.execute_reply.started":"2026-02-03T13:38:47.895758Z","shell.execute_reply":"2026-02-03T13:38:47.907767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DropPath(nn.Module):\n    def __init__(self, drop_prob=0.0):\n        super().__init__()\n        self.drop_prob = drop_prob\n\n    def forward(self, x):\n        if self.drop_prob == 0. or not self.training:\n            return x\n\n        keep_prob = 1 - self.drop_prob\n        shape = (x.shape[0],) + (1,) * (x.ndim - 1)\n        random_tensor = x.new_empty(shape).bernoulli_(keep_prob)\n        if keep_prob > 0.0:\n            random_tensor.div_(keep_prob)\n        return x * random_tensor\n\nclass FeedForward(nn.Module):\n    def __init__(self, embed_dim, hidden_dim=None, dropout=0.0, bias=False, use_glu=True, use_checkpoint=False):\n        super().__init__()\n\n        if hidden_dim is None:\n            hidden_dim = int(8 * embed_dim / 3) if use_glu else 4 * embed_dim\n            hidden_dim = ((hidden_dim + 255) // 256) * 256\n\n        self.use_glu = use_glu\n        self.use_checkpoint = use_checkpoint\n\n        if use_glu:\n            self.w12 = nn.Linear(embed_dim, 2 * hidden_dim, bias=bias)\n        else:\n            self.w1 = nn.Linear(embed_dim, hidden_dim, bias=bias)\n\n        self.w_out = nn.Linear(hidden_dim, embed_dim, bias=bias)\n        self.dropout = nn.Dropout(dropout)\n\n    def _compute(self, x):\n        if self.use_glu:\n            x12 = self.w12(x)\n            gate, value = x12.chunk(2, dim=-1)\n            return self.w_out(F.silu(gate) * value)\n        else:\n            return self.w_out(F.gelu(self.w1(x)))\n\n    def forward(self, x):\n        if self.use_checkpoint and self.training:\n            if x.requires_grad:\n                out = checkpoint(self._compute, x, use_reentrant=False)\n            else:\n                out = self._compute(x)\n        else:\n            out = self._compute(x)\n\n        return self.dropout(out)\n\nclass MoEFeedForward(nn.Module):\n    def __init__(self, num_experts, top_k, embed_dim, capacity_factor=1.0, **kwargs):\n        super().__init__()\n        self.num_experts = num_experts\n        self.top_k = top_k\n        self.embed_dim = embed_dim\n        self.capacity_factor = capacity_factor\n\n        self.router = nn.Linear(embed_dim, num_experts, bias=False)\n        self.experts = nn.ModuleList([\n            FeedForward(embed_dim, **kwargs) for _ in range(num_experts)\n        ])\n\n        self.register_buffer(\"capacity\", torch.tensor(0, dtype=torch.long), persistent=False)\n        self.last_num_tokens = -1\n\n    def update_capacity(self, num_tokens):\n        if num_tokens == self.last_num_tokens:\n            return\n        cap = int(math.ceil((num_tokens * self.top_k / self.num_experts) * self.capacity_factor))\n        align = 16\n        if cap % align != 0:\n            cap += (align - (cap % align))\n        self.capacity.fill_(cap)\n        self.last_num_tokens = num_tokens\n\n    def forward(self, x):\n        original_shape = x.shape\n        batch_size, seq_len, dim = x.shape\n        num_tokens = batch_size * seq_len\n        x_flat = x.view(-1, dim)\n\n        if num_tokens != self.last_num_tokens:\n            self.update_capacity(num_tokens)\n        capacity = self.capacity.item()\n\n        gate_logits = self.router(x_flat)\n        routing_weights = F.softmax(gate_logits, dim=-1)\n\n        weights, indices = torch.topk(routing_weights, self.top_k, dim=-1)\n        weights = weights / weights.sum(dim=-1, keepdim=True)\n\n        flat_weights = weights.view(-1)\n        flat_indices = indices.view(-1)\n\n        sorted_weights, sorted_args = torch.sort(flat_weights, descending=True)\n        sorted_expert_indices = flat_indices[sorted_args]\n\n        mask_sorted = F.one_hot(sorted_expert_indices, num_classes=self.num_experts).int()\n        position_in_expert_sorted = torch.cumsum(mask_sorted, dim=0) - 1\n        token_priority_sorted = (position_in_expert_sorted * mask_sorted).sum(dim=-1)\n\n        token_priority = torch.zeros_like(token_priority_sorted)\n        token_priority[sorted_args] = token_priority_sorted\n\n        valid_mask = token_priority < capacity\n\n        dest_indices = flat_indices * capacity + token_priority\n        dest_indices = dest_indices * valid_mask\n\n        expert_input_flat = torch.zeros(\n            self.num_experts * capacity, dim, dtype=x.dtype, device=x.device\n        )\n\n        x_repeated = x_flat.repeat_interleave(self.top_k, dim=0)\n        x_repeated_valid = x_repeated * valid_mask.unsqueeze(-1)\n\n        expert_input_flat.index_add_(0, dest_indices, x_repeated_valid)\n        expert_input_split = expert_input_flat.view(self.num_experts, capacity, dim)\n\n        expert_outputs = []\n        for i, expert in enumerate(self.experts):\n            out = expert(expert_input_split[i])\n            expert_outputs.append(out)\n\n        expert_output_flat = torch.stack(expert_outputs).view(-1, dim)\n\n        output_gathered = torch.index_select(expert_output_flat, 0, dest_indices)\n        output_gathered = output_gathered * valid_mask.unsqueeze(-1)\n        output_weighted = output_gathered * weights.view(-1, 1)\n\n        final_output = output_weighted.view(num_tokens, self.top_k, dim).sum(dim=1)\n        return final_output.view(original_shape)\n\n\nclass RotaryPositionalEmbedding(nn.Module):\n    def __init__(self, dim, max_seq_len=2048, base=10000.0, interleaved=False, device=None):\n        super().__init__()\n        self.dim = dim\n        self.max_seq_len = max_seq_len\n        self.base = base\n        self.interleaved = interleaved\n\n        inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32, device=device) / dim))\n        self.register_buffer(\"inv_freq\", inv_freq, persistent=False)\n\n        self._seq_len_cached = 0\n        self.register_buffer(\"cos_cached\", None, persistent=False)\n        self.register_buffer(\"sin_cached\", None, persistent=False)\n\n        if device is not None:\n            self._build_cache(max_seq_len, device=device)\n\n    def _build_cache(self, seq_len, device=None):\n        if seq_len > self.max_seq_len:\n            self.max_seq_len = seq_len\n\n        t = torch.arange(seq_len, device=device, dtype=self.inv_freq.dtype)\n        freqs = torch.outer(t, self.inv_freq)\n\n        if not self.interleaved:\n            emb = torch.cat((freqs, freqs), dim=-1)\n        else:\n            emb = torch.repeat_interleave(freqs, 2, dim=-1)\n\n        self.register_buffer(\"cos_cached\", emb.cos(), persistent=False)\n        self.register_buffer(\"sin_cached\", emb.sin(), persistent=False)\n        self._seq_len_cached = seq_len\n\n    def _rotate_half(self, x):\n        x1, x2 = x.chunk(2, dim=-1)\n        return torch.cat((-x2, x1), dim=-1)\n\n    def _rotate_interleaved(self, x):\n        x_reshaped = x.view(*x.shape[:-1], -1, 2)\n        x_rotated = torch.stack((-x_reshaped[..., 1], x_reshaped[..., 0]), dim=-1)\n        return x_rotated.flatten(-2)\n\n    def forward(self, x, seq_len=None, offset=0):\n        if seq_len is None:\n            seq_len = x.shape[-2]\n\n        needed = offset + seq_len\n\n        if needed > self._seq_len_cached:\n            self._build_cache(max(needed, self.max_seq_len), device=x.device)\n\n        cos = self.cos_cached[offset:offset + seq_len].to(dtype=x.dtype)[None, None, :, :]\n        sin = self.sin_cached[offset:offset + seq_len].to(dtype=x.dtype)[None, None, :, :]\n\n        if not self.interleaved:\n            return (x * cos) + (self._rotate_half(x) * sin)\n        else:\n            return (x * cos) + (self._rotate_interleaved(x) * sin)\n\n\nclass FlashGQA(nn.Module):\n    def __init__(\n        self,\n        embed_dim,\n        num_heads_q,\n        num_heads_kv,\n        rope=None,\n        dropout=0.0,\n        bias=False,\n        is_cross_attention=False,\n        context_dim=None,\n        norm_eps=1e-6\n    ):\n        super().__init__()\n\n        if embed_dim % num_heads_q != 0:\n            raise ValueError(f\"embed_dim {embed_dim} must be divisible by num_heads_q\")\n\n        self.num_heads_q = num_heads_q\n        self.num_heads_kv = num_heads_kv\n        self.head_dim = embed_dim // num_heads_q\n        self.is_cross_attention = is_cross_attention\n        init_scale = self.head_dim ** -0.5\n        self.scale = nn.Parameter(torch.tensor([init_scale]))\n        self.rep = num_heads_q // num_heads_kv\n\n        if is_cross_attention:\n            self.rope = None\n        else:\n            self.rope = rope\n\n        if not is_cross_attention:\n            total_dim = (num_heads_q + 2 * num_heads_kv) * self.head_dim\n            self.qkv_proj = nn.Linear(embed_dim, total_dim, bias=bias)\n        else:\n            ctx_dim = context_dim or embed_dim\n            self.q_proj = nn.Linear(embed_dim, num_heads_q * self.head_dim, bias=bias)\n            self.kv_proj = nn.Linear(ctx_dim, 2 * num_heads_kv * self.head_dim, bias=bias)\n\n        self.eps = norm_eps\n        self.out_proj = nn.Linear(num_heads_q * self.head_dim, embed_dim, bias=bias)\n        self.dropout_p = dropout\n\n    def _add_causal_mask(self, mask, seq_len):\n        dtype = mask.dtype\n        device = mask.device\n        causal_mask = torch.full(\n            (seq_len, seq_len),\n            torch.finfo(dtype).min,\n            device=device,\n            dtype=dtype\n        )\n        causal_mask = torch.triu(causal_mask, diagonal=1)\n\n        return causal_mask.unsqueeze(0).unsqueeze(0) + mask\n    def forward(self, x,\n                context=None,\n                attn_mask=None,\n                is_causal=False,\n                offset=0,\n                kv_cache=None,\n                use_cache=False):\n        B, Lq, _ = x.shape\n\n        if not self.is_cross_attention:\n            qkv = self.qkv_proj(x)\n            q, k, v = qkv.split([\n                self.num_heads_q * self.head_dim,\n                self.num_heads_kv * self.head_dim,\n                self.num_heads_kv * self.head_dim\n            ], dim=-1)\n\n            q = q.view(B, Lq, self.num_heads_q, self.head_dim).transpose(1, 2)\n            k = k.view(B, Lq, self.num_heads_kv, self.head_dim).transpose(1, 2)\n            v = v.view(B, Lq, self.num_heads_kv, self.head_dim).transpose(1, 2)\n\n            if self.rope is not None:\n                q = self.rope(q, offset=offset)\n                k = self.rope(k, offset=offset)\n\n            if use_cache and kv_cache is not None:\n                past_k, past_v = kv_cache\n                k = torch.cat([past_k, k], dim=2)\n                v = torch.cat([past_v, v], dim=2)\n            current_kv_cache = (k.detach(), v.detach()) if use_cache else None\n\n        else:\n            is_causal = False\n            if context.dim() == 4:\n                context = context.flatten(1, 2)\n\n            q = self.q_proj(x)\n            q = q.view(B, Lq, self.num_heads_q, self.head_dim).transpose(1, 2)\n            if use_cache and kv_cache is not None:\n                k, v = kv_cache\n            else:\n                if context.dim() == 4:\n                    context = context.flatten(1, 2)\n                kv = self.kv_proj(context)\n                k, v = kv.split([self.num_heads_kv * self.head_dim] * 2, dim=-1)\n\n                Lk = context.shape[1]\n                k = k.view(B, Lk, self.num_heads_kv, self.head_dim).transpose(1, 2)\n                v = v.view(B, Lk, self.num_heads_kv, self.head_dim).transpose(1, 2)\n\n            current_kv_cache = (k.detach(), v.detach()) if use_cache else None\n        \n        q = F.normalize(q, p=2, dim=-1, eps=self.eps)\n        k = F.normalize(k, p=2, dim=-1, eps=self.eps)\n        \n        if self.rep > 1:\n            k = k[:, :, None, :, :].expand(B, self.num_heads_kv, self.rep, -1, -1)\n            k = k.reshape(B, self.num_heads_q, -1, self.head_dim)\n            v = v[:, :, None, :, :].expand(B, self.num_heads_kv, self.rep, -1, -1)\n            v = v.reshape(B, self.num_heads_q, -1, self.head_dim)\n\n        actual_is_causal = is_causal\n        if kv_cache is not None:\n            actual_is_causal = False\n        if attn_mask is not None:\n            if attn_mask.dim() == 2:\n                attn_mask = attn_mask.view(B, 1, 1, -1)\n            if is_causal and kv_cache is None:\n                attn_mask = self._add_causal_mask(attn_mask, Lq)\n                actual_is_causal = False\n            else:\n                actual_is_causal = False\n\n        q = q * self.scale\n        out = F.scaled_dot_product_attention(\n            q, k, v, \n            attn_mask=attn_mask,\n            dropout_p=self.dropout_p if self.training else 0.0,\n            is_causal=actual_is_causal, \n            scale=1.0\n        )\n\n        out = out.transpose(1, 2).contiguous().view(B, Lq, -1)\n        out = self.out_proj(out)\n\n        if use_cache:\n            return out, current_kv_cache\n        else:\n            return out\n\n\nclass TransformerBlock(nn.Module):\n    def __init__(\n        self,\n        embed_dim,\n        num_heads_q,\n        num_heads_kv=None,\n        hidden_dim=None,\n        dropout=0.0,\n        drop_path=0.0,\n        max_seq_len=2048,\n        rope_base=10000.0,\n        rope_interleaved=False,\n        rope=None,\n        use_glu=True,\n        has_cross_attention=False,\n        context_dim=None,\n        bias=False,\n        norm_eps=1e-6,\n        use_checkpoint=False\n    ):\n        super().__init__()\n        if num_heads_kv is None:\n            num_heads_kv = num_heads_q\n\n        if rope is None:\n            self.rope = RotaryPositionalEmbedding(\n                dim=embed_dim // num_heads_q,\n                max_seq_len=max_seq_len,\n                base=rope_base,\n                interleaved=rope_interleaved\n            )\n        else:\n            self.rope = rope\n\n        self.norm1 = nn.RMSNorm(embed_dim, eps=norm_eps)\n\n        self.self_attn = FlashGQA(\n            embed_dim=embed_dim,\n            num_heads_q=num_heads_q,\n            num_heads_kv=num_heads_kv,\n            dropout=dropout,\n            bias=bias,\n            rope=self.rope,\n            is_cross_attention=False,\n            norm_eps=norm_eps\n        )\n\n        self.has_cross_attention = has_cross_attention\n        if has_cross_attention:\n            self.norm_cross = nn.RMSNorm(embed_dim, eps=norm_eps)\n            self.cross_attn = FlashGQA(\n                embed_dim=embed_dim,\n                num_heads_q=num_heads_q,\n                num_heads_kv=num_heads_kv,\n                dropout=dropout,\n                bias=bias,\n                rope=None,\n                is_cross_attention=True,\n                context_dim=context_dim,\n                norm_eps=norm_eps\n            )\n\n        self.norm2 = nn.RMSNorm(embed_dim, eps=norm_eps)\n\n        self.ffn = FeedForward(\n            embed_dim=embed_dim,\n            hidden_dim=hidden_dim,\n            dropout=dropout,\n            bias=bias,\n            use_glu=use_glu,\n            use_checkpoint=use_checkpoint\n        )\n\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n\n    def forward(self, x,\n                context=None,\n                mask=None,\n                cross_mask=None,\n                is_causal=True,\n                offset=0,\n                use_cache=False,\n                kv_cache_attn=None,\n                kv_cache_cross=None):\n        new_kv_cache_attn = None\n        new_kv_cache_cross = None\n        if use_cache:\n            out, new_kv_cache_attn = self.self_attn(\n                self.norm1(x),\n                attn_mask=mask,\n                is_causal=is_causal,\n                offset=offset,\n                use_cache=use_cache,\n                kv_cache=kv_cache_attn\n            )\n            x = x + self.drop_path(out)\n        else:\n            shortcut = x\n            x = self.self_attn(\n                self.norm1(x),\n                attn_mask=mask,\n                is_causal=is_causal,\n                offset=offset\n            )\n            x = shortcut + self.drop_path(x)\n\n        if self.has_cross_attention and context is not None:\n            if use_cache:\n                out, new_kv_cache_cross = self.cross_attn(\n                    self.norm_cross(x),\n                    context=context,\n                    attn_mask=cross_mask,\n                    is_causal=False,\n                    use_cache=use_cache,\n                    kv_cache=kv_cache_cross\n                )\n                x = x + self.drop_path(out)\n            else:\n                shortcut = x\n                x = self.cross_attn(\n                    self.norm_cross(x),\n                    context=context,\n                    attn_mask=cross_mask,\n                    is_causal=False,\n                    use_cache=use_cache,\n                    kv_cache=kv_cache_cross\n                )\n                x = shortcut + self.drop_path(x)\n        shortcut = x\n        x = self.ffn(self.norm2(x))\n        x = shortcut + self.drop_path(x)\n        if use_cache:\n            if self.has_cross_attention:\n                return x, (new_kv_cache_attn, new_kv_cache_cross)\n            else:\n                return x, new_kv_cache_attn\n        else:\n            return x\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:38:48.04449Z","iopub.execute_input":"2026-02-03T13:38:48.044796Z","iopub.status.idle":"2026-02-03T13:38:48.08863Z","shell.execute_reply.started":"2026-02-03T13:38:48.044772Z","shell.execute_reply":"2026-02-03T13:38:48.087683Z"},"id":"MFX6O-u8p2Vx"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Seq2SeqLLM(nn.Module):\n    def __init__(\n        self,\n        model_name=\"answerdotai/ModernBERT-base\",\n        num_block=12,\n        num_heads_q=12,\n        num_heads_kv=4,\n        hidden_dim=None,\n        dropout=0.0,\n        drop_path=0.2,\n        max_seq_len=128,\n        use_checkpoint=False,\n    ):\n        super().__init__()\n\n        source_model = AutoModelForMaskedLM.from_pretrained(\n            model_name,\n            torch_dtype=torch.float16,\n            device_map=\"cpu\"\n        )\n        pretrained_embedding = source_model.get_input_embeddings().weight.detach()\n\n        self.vocab_size, self.embed_dim = pretrained_embedding.shape\n        self.embeddings = nn.Embedding.from_pretrained(\n            pretrained_embedding,\n            freeze=False\n        ).to(torch.float32)\n\n        dpr = [x.item() for x in torch.linspace(0, drop_path, num_block)]\n\n        self.encoder_norm = nn.RMSNorm(self.embed_dim, eps=1e-5)\n        self.encoder_blocks = nn.ModuleList([\n            TransformerBlock(\n                embed_dim=self.embed_dim,\n                num_heads_q=num_heads_q,\n                num_heads_kv=num_heads_kv,\n                hidden_dim=hidden_dim,\n                dropout=dropout,\n                drop_path=dpr[i],\n                max_seq_len=max_seq_len,\n                rope_base=20000.0,\n                use_glu=True,\n                has_cross_attention=False,\n                norm_eps=1e-5,\n                use_checkpoint=use_checkpoint\n            ) for i in range(num_block // 2)\n        ])\n\n        self.decoder_norm = nn.RMSNorm(self.embed_dim, eps=1e-5)\n        self.decoder_blocks = nn.ModuleList([\n            TransformerBlock(\n                embed_dim=self.embed_dim,\n                num_heads_q=num_heads_q,\n                num_heads_kv=num_heads_kv,\n                hidden_dim=hidden_dim,\n                dropout=dropout,\n                drop_path=dpr[i + num_block // 2],\n                max_seq_len=max_seq_len,\n                rope_base=20000.0,\n                use_glu=True,\n                has_cross_attention=True,\n                context_dim=self.embed_dim,\n                norm_eps=1e-5,\n                use_checkpoint=use_checkpoint\n            ) for i in range(num_block // 2)\n        ])\n\n        self.lm_head = nn.Linear(self.embed_dim, self.vocab_size, bias=False)\n        self.lm_head.weight = self.embeddings.weight\n\n    def create_mask(self, mask, dtype):\n        if mask is None:\n            return None\n        inverted_mask = 1.0 - mask.to(dtype)\n        min_dtype = torch.finfo(dtype).min\n        mask_for_attn = inverted_mask * min_dtype\n        if mask_for_attn.dim() == 2:\n            mask_for_attn = mask_for_attn.unsqueeze(1).unsqueeze(2)\n        return mask_for_attn\n\n    def forward_encoder(self, input_ids, attention_mask):\n        x = self.embeddings(input_ids)\n        mask_for_attn = self.create_mask(attention_mask, x.dtype)\n        for layer in self.encoder_blocks:\n            x = layer(x, mask=mask_for_attn, is_causal=False)\n        x = self.encoder_norm(x)\n        return x\n\n    def forward_decoder(\n        self,\n        decoder_input_ids,\n        decoder_attention_mask,\n        encoder_hidden_states,\n        encoder_attention_mask,\n        kv_caches=None,\n        use_cache=False\n    ):\n        x = self.embeddings(decoder_input_ids)\n\n        dec_mask_for_attn = self.create_mask(decoder_attention_mask, x.dtype)\n        cross_mask_for_attn = self.create_mask(encoder_attention_mask, x.dtype)\n\n        new_kv_caches = []\n\n        for i, layer in enumerate(self.decoder_blocks):\n            layer_cache_attn = kv_caches[i][0] if kv_caches else None\n            layer_cache_cross = kv_caches[i][1] if kv_caches else None\n\n            if use_cache:\n                x, (new_k_attn, new_k_cross) = layer(\n                    x,\n                    context=encoder_hidden_states,\n                    mask=dec_mask_for_attn,\n                    cross_mask=cross_mask_for_attn,\n                    is_causal=True,\n                    use_cache=True,\n                    kv_cache_attn=layer_cache_attn,\n                    kv_cache_cross=layer_cache_cross\n                )\n                new_kv_caches.append((new_k_attn, new_k_cross))\n            else:\n                x = layer(\n                    x,\n                    context=encoder_hidden_states,\n                    mask=dec_mask_for_attn,\n                    cross_mask=cross_mask_for_attn,\n                    is_causal=True,\n                    use_cache=False\n                )\n\n        x = self.decoder_norm(x)\n        logits = self.lm_head(x)\n\n        if use_cache:\n            return logits, new_kv_caches\n        return logits\n\n    def forward(\n        self,\n        input_ids,\n        attention_mask,\n        decoder_input_ids,\n        decoder_attention_mask,\n        **kwargs\n    ):\n        encoder_output = self.forward_encoder(input_ids, attention_mask)\n        logits = self.forward_decoder(\n            decoder_input_ids=decoder_input_ids,\n            decoder_attention_mask=decoder_attention_mask,\n            encoder_hidden_states=encoder_output,\n            encoder_attention_mask=attention_mask\n        )\n        return logits\n\n    @torch.no_grad()\n    def generate(\n        self,\n        tokenizer,\n        text: str,\n        max_new_tokens: int = 128,\n        device: str = \"cuda\",\n        temperature: float = 0.5,\n        top_k: int = 5,\n        repetition_penalty: float = 1.0,\n        verbose: bool = False, \n    ):\n        self.eval()\n        self.to(device)\n        \n        if verbose:\n            print(f\"[-] Input: {text}\")\n\n        eos_token_id = tokenizer.sep_token_id\n        bos_token_id = tokenizer.cls_token_id\n\n        raw_input_ids = tokenizer.encode(text, add_special_tokens=False)\n        input_ids_list = [bos_token_id] + raw_input_ids + [eos_token_id]\n        \n        input_ids = torch.tensor([input_ids_list], device=device)\n        attention_mask = torch.ones_like(input_ids).to(device)\n\n        encoder_hidden_states = self.forward_encoder(input_ids, attention_mask)\n\n        decoder_input_ids = torch.tensor([[bos_token_id]], device=device)\n        kv_caches = None\n\n        for _ in range(max_new_tokens):\n            if kv_caches is None:\n                curr_input = decoder_input_ids\n            else:\n                curr_input = decoder_input_ids[:, -1:]\n\n            logits, kv_caches = self.forward_decoder(\n                decoder_input_ids=curr_input,\n                decoder_attention_mask=None,\n                encoder_hidden_states=encoder_hidden_states,\n                encoder_attention_mask=attention_mask,\n                kv_caches=kv_caches,\n                use_cache=True\n            )\n\n            next_token_logits = logits[:, -1, :]\n\n            if repetition_penalty != 1.0:\n                generated_tokens = decoder_input_ids[0].unique()\n                score = next_token_logits[:, generated_tokens]\n                score = torch.where(\n                    score < 0, \n                    score * repetition_penalty, \n                    score / repetition_penalty\n                )\n                next_token_logits[:, generated_tokens] = score\n\n            if temperature > 0:\n                next_token_logits = next_token_logits / temperature\n\n            if top_k is not None:\n                v, _ = torch.topk(\n                    next_token_logits, \n                    min(top_k, next_token_logits.size(-1))\n                )\n                next_token_logits[\n                    next_token_logits < v[:, [-1]]\n                ] = -float(\"Inf\")\n\n            probs = F.softmax(next_token_logits, dim=-1)\n            next_token = torch.multinomial(probs, num_samples=1)\n\n            decoder_input_ids = torch.cat(\n                [decoder_input_ids, next_token], \n                dim=1\n            )\n\n            if next_token.item() == tokenizer.sep_token_id:\n                break\n\n        output_text = tokenizer.decode(\n            decoder_input_ids[0], \n            skip_special_tokens=True\n        )\n\n        if verbose:\n            print(f\"[-] Result: {output_text}\")\n            print(\"-\" * 50)\n\n        return output_text\nmodel = Seq2SeqLLM()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:39:23.019433Z","iopub.execute_input":"2026-02-03T13:39:23.020097Z","iopub.status.idle":"2026-02-03T13:39:24.506884Z","shell.execute_reply.started":"2026-02-03T13:39:23.020071Z","shell.execute_reply":"2026-02-03T13:39:24.506238Z"},"id":"RNI8blprp2Vz","outputId":"927f4a94-a1ee-4621-e991-b9caeb50fd51"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport random\nimport numpy as np\nimport os\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    print(f\"Global seed set to {seed}\")\n\nseed_everything(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:38:49.770352Z","iopub.execute_input":"2026-02-03T13:38:49.770649Z","iopub.status.idle":"2026-02-03T13:38:49.777994Z","shell.execute_reply.started":"2026-02-03T13:38:49.770627Z","shell.execute_reply":"2026-02-03T13:38:49.777177Z"},"id":"NqVGSrwQp2Vy","outputId":"f8b12123-1a60-4b38-ad9b-6a614f3a17e0"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"try:\n    nltk.data.find('tokenizers/punkt')\nexcept LookupError:\n    nltk.download('punkt')\n    nltk.download('punkt_tab')\n\nclass MRPCSDataset(Dataset):\n    def __init__(\n        self,\n        model_name=\"answerdotai/ModernBERT-base\",\n        dataset_name=(\"glue\", \"qqp\"),\n        split=\"train\",\n        max_len=128,\n        num_sample=None\n    ):\n        self.tokenizer = AutoTokenizer.from_pretrained(model_name)\n\n        self.bos_token_id = self.tokenizer.cls_token_id\n        self.eos_token_id = self.tokenizer.sep_token_id\n        self.pad_token_id = self.tokenizer.pad_token_id\n        print(self.bos_token_id)\n        print(self.eos_token_id)\n        print(self.pad_token_id)\n        \n        if self.tokenizer.pad_token is None:\n            self.tokenizer.pad_token = self.tokenizer.eos_token\n\n        self.max_len = max_len\n        self.samples = []\n\n        path, name = dataset_name\n        dataset = load_dataset(path, name, split=split)\n\n        self._build_dataset(dataset, num_sample)\n        print(f\"Build dataset complete, found {len(self.samples)} samples\")\n        \n    def clean_text(self, text):\n        if not text:\n            return \"\"\n        return re.sub(r'\\s+', ' ', text).strip()\n\n    def _build_dataset(self, dataset, num_sample):\n        k1, k2 = ('question1', 'question2') if 'question1' in dataset.column_names else ('sentence1', 'sentence2')\n\n        for item in tqdm(dataset):\n            if item['label'] == 1:\n                s1 = self.clean_text(item[k1])\n                s2 = self.clean_text(item[k2])\n\n                if s1 and s2:\n                    self.create_sample(s1, s2)\n                    self.create_sample(s2, s1)\n\n            if num_sample is not None and len(self.samples) > num_sample:\n                break\n\n    def create_sample(self, source, target):\n        src_tokens = self.tokenizer.encode(source, add_special_tokens=False)\n\n        if len(src_tokens) > self.max_len - 2:\n            src_tokens = src_tokens[:self.max_len - 2]\n\n        src_ids = [self.bos_token_id] + src_tokens + [self.eos_token_id]\n\n        pad_len = self.max_len - len(src_ids)\n        input_ids = src_ids + [self.pad_token_id] * pad_len\n        attention_mask = [1] * len(src_ids) + [0] * pad_len\n\n        tgt_tokens = self.tokenizer.encode(target, add_special_tokens=False)\n\n        if len(tgt_tokens) > self.max_len - 2:\n            tgt_tokens = tgt_tokens[:self.max_len - 2]\n\n        full_tgt_ids = [self.bos_token_id] + tgt_tokens + [self.eos_token_id]\n\n        pad_len_dec = self.max_len - len(full_tgt_ids)\n\n        decoder_input_ids = full_tgt_ids + [self.pad_token_id] * pad_len_dec\n        decoder_attention_mask = [1] * len(full_tgt_ids) + [0] * pad_len_dec\n\n        labels = full_tgt_ids + [-100] * pad_len_dec\n\n        sample = {\n            \"input_ids\": torch.tensor(input_ids, dtype=torch.long),\n            \"attention_mask\": torch.tensor(attention_mask, dtype=torch.long),\n            \"decoder_input_ids\": torch.tensor(decoder_input_ids, dtype=torch.long),\n            \"decoder_attention_mask\": torch.tensor(decoder_attention_mask, dtype=torch.long),\n            \"labels\": torch.tensor(labels, dtype=torch.long),\n        }\n\n        self.samples.append(sample)\n        return sample\n\n\n    \n        self.samples.append(sample)\n        return sample\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        return self.samples[idx]\n\n\ndef collate_fn(batch):\n    batch_output = {}\n    keys = batch[0].keys() \n    for key in keys:\n        batch_output[key] = torch.stack([item[key] for item in batch])\n        \n    return batch_output\n\ndtrain = MRPCSDataset(split=\"train\", num_sample=None)\ndtest = MRPCSDataset(split=\"validation\", num_sample=None)\n\nbatch_size = 64\ndloader_train = DataLoader(dtrain, batch_size=batch_size, shuffle=True, num_workers=2, collate_fn=collate_fn, pin_memory=True)\ndloader_test = DataLoader(dtest, batch_size=batch_size, shuffle=False, num_workers=2, collate_fn=collate_fn, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:38:49.778872Z","iopub.execute_input":"2026-02-03T13:38:49.779642Z","iopub.status.idle":"2026-02-03T13:38:58.67911Z","shell.execute_reply.started":"2026-02-03T13:38:49.779614Z","shell.execute_reply":"2026-02-03T13:38:58.677485Z"},"id":"V-QKsJuep2V0"},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Trainer:\n    def __init__(\n        self,\n        model,\n        train_loader,\n        val_loader,\n        optimizer,\n        scheduler,\n        criterion,\n        metric=None,\n        mode='min',\n        device='cuda',\n        save_dir='./checkpoints',\n        hf_token=None,\n        hf_repo_id=None,\n        use_flash=True,\n        early_stopping_patience=5,\n        early_stopping_threshold=0.0,\n        early_stopping_min_epochs=20,\n        gradient_accumulation_steps=1,\n        max_grad_norm=1.0\n    ):\n        self.device = device\n        self.model = model.to(device)\n\n        self.model.apply(self._init_weights)\n        print(\"-> Weights Initialized (Safe Mode).\")\n\n        if use_flash:\n            print(\"-> Compiling model with torch.compile...\")\n            try:\n                self.model = torch.compile(self.model)\n                print(\"   Compile Success!\")\n            except Exception as e:\n                print(f\"   Compile Failed/Skipped: {e}\")\n\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.criterion = criterion\n        self.metric = metric\n        self.mode = mode\n\n        self.scaler = torch.amp.GradScaler('cuda')\n        self.save_dir = save_dir\n\n        if not os.path.exists(save_dir):\n            os.makedirs(save_dir)\n\n        self.hf_repo_id = hf_repo_id\n        self.hf_api = None\n        self.hf_token = hf_token\n        if hf_token and hf_repo_id:\n            print(f\"-> Login to hugging face to push: {hf_repo_id}\")\n            login(token=hf_token)\n            self.hf_api = HfApi()\n            self.hf_api.create_repo(repo_id=hf_repo_id, exist_ok=True)\n        else:\n            print(f\"-> Skip HF Push (Token or Repo ID missing)\")\n\n        if mode == 'max':\n            self.best_metric = -float('inf')\n        else:\n            self.best_metric = float('inf')\n\n        self.patience = early_stopping_patience\n        self.threshold = early_stopping_threshold\n        self.min_epochs = early_stopping_min_epochs\n        self.grad_accum_steps = gradient_accumulation_steps\n        self.max_grad_norm = max_grad_norm\n\n        self.patience_counter = 0\n        self.early_stop = False\n\n    def _get_vram(self):\n        if torch.cuda.is_available():\n            return torch.cuda.memory_allocated() / 1024**3\n        return 0.0\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Embedding):\n            return\n\n        if isinstance(m, nn.Linear):\n            if hasattr(self, 'model') and hasattr(self.model, 'embeddings'):\n                if hasattr(self.model, 'embeddings') and \\\n                   m.weight.data_ptr() == self.model.embeddings.weight.data_ptr():\n                    return\n\n            nn.init.trunc_normal_(m.weight, std=0.02)\n            if m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.weight, 1.0)\n            nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.Conv2d):\n            nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            if m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n\n    def train_one_epoch(self, epoch_idx):\n            self.model.train()\n            running_loss = 0.0\n            total_samples = 0\n            start_time = time.time()\n    \n            pbar = tqdm(self.train_loader, desc=f\"Epoch {epoch_idx+1} [Train]\", leave=False)\n    \n            self.optimizer.zero_grad(set_to_none=True)\n    \n            for i, batch in enumerate(pbar):\n                input_ids = batch['input_ids'].to(self.device, non_blocking=True)\n                attention_mask = batch['attention_mask'].to(self.device, non_blocking=True)\n    \n                dec_input_ids = batch['decoder_input_ids'].to(self.device, non_blocking=True)\n                dec_attention_mask = batch['decoder_attention_mask'].to(self.device, non_blocking=True)\n                labels = batch['labels'].to(self.device, non_blocking=True)\n    \n                batch_size = input_ids.size(0)\n    \n                model_dec_input = dec_input_ids[:, :-1]\n                model_dec_mask = dec_attention_mask[:, :-1]\n                loss_target = labels[:, 1:]\n    \n                with torch.amp.autocast('cuda'):\n                    logits = self.model(\n                        input_ids=input_ids,\n                        attention_mask=attention_mask,\n                        decoder_input_ids=model_dec_input,\n                        decoder_attention_mask=model_dec_mask\n                    )\n    \n                    loss = self.criterion(\n                        logits.reshape(-1, logits.size(-1)),\n                        loss_target.reshape(-1)\n                    )\n    \n                loss_scaled = loss / self.grad_accum_steps\n                self.scaler.scale(loss_scaled).backward()\n    \n                if (i + 1) % self.grad_accum_steps == 0 or (i + 1) == len(self.train_loader):\n                    self.scaler.unscale_(self.optimizer)\n                    torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=self.max_grad_norm)\n    \n                    self.scaler.step(self.optimizer)\n                    self.scaler.update()\n                    self.optimizer.zero_grad(set_to_none=True)\n    \n                loss_val = loss.item()\n                running_loss += loss_val * batch_size\n                total_samples += batch_size\n    \n                vram = self._get_vram()\n                postfix_dict = {'Loss': f\"{loss_val:.4f}\", 'VRAM': f'{vram:.2f}GB'}\n    \n                if self.metric:\n                    try:\n                        postfix_dict['PPL'] = f\"{self.metric(loss_val):.2f}\"\n                    except OverflowError:\n                        postfix_dict['PPL'] = \"inf\"\n    \n                pbar.set_postfix(postfix_dict)\n    \n            epoch_loss = running_loss / total_samples\n            duration = time.time() - start_time\n    \n            epoch_metric = 0.0\n            if self.metric:\n                try:\n                    epoch_metric = self.metric(epoch_loss)\n                except:\n                    epoch_metric = float('inf')\n    \n            return epoch_loss, epoch_metric, duration\n    \n    \n    @torch.no_grad()\n    def evaluate(self, loader, mode=\"Val\"):\n            self.model.eval()\n            running_loss = 0.0\n            total_samples = 0\n    \n            pbar = tqdm(loader, desc=f\"Evaluating [{mode}]\", leave=False)\n    \n            for batch in pbar:\n                input_ids = batch['input_ids'].to(self.device, non_blocking=True)\n                attention_mask = batch['attention_mask'].to(self.device, non_blocking=True)\n    \n                dec_input_ids = batch['decoder_input_ids'].to(self.device, non_blocking=True)\n                dec_attention_mask = batch['decoder_attention_mask'].to(self.device, non_blocking=True)\n                labels = batch['labels'].to(self.device, non_blocking=True)\n    \n                batch_size = input_ids.size(0)\n    \n                model_dec_input = dec_input_ids[:, :-1]\n                model_dec_mask = dec_attention_mask[:, :-1]\n                loss_target = labels[:, 1:]\n    \n                with torch.amp.autocast('cuda'):\n                    logits = self.model(\n                        input_ids=input_ids,\n                        attention_mask=attention_mask,\n                        decoder_input_ids=model_dec_input,\n                        decoder_attention_mask=model_dec_mask\n                    )\n    \n                    loss = self.criterion(\n                        logits.reshape(-1, logits.size(-1)),\n                        loss_target.reshape(-1)\n                    )\n    \n                loss_val = loss.item()\n                running_loss += loss_val * batch_size\n                total_samples += batch_size\n    \n            epoch_loss = running_loss / total_samples\n    \n            epoch_metric = 0.0\n            if self.metric:\n                try:\n                    epoch_metric = self.metric(epoch_loss)\n                except:\n                    epoch_metric = float('inf')\n    \n            return epoch_loss, epoch_metric\n\n\n    def save_checkpoint(self, epoch, filename='last.pt'):\n        filepath = os.path.join(self.save_dir, filename)\n        model_to_save = self.model.module if hasattr(self.model, \"module\") else self.model\n        model_to_save = model_to_save._orig_mod if hasattr(model_to_save, \"_orig_mod\") else model_to_save\n\n        checkpoint = {\n            'epoch': epoch,\n            'state_dict': model_to_save.state_dict(),\n            'optimizer': self.optimizer.state_dict(),\n            'scaler': self.scaler.state_dict(),\n            'scheduler': self.scheduler.state_dict() if self.scheduler else None,\n            'best_metric': self.best_metric,\n        }\n        torch.save(checkpoint, filepath)\n\n        if self.hf_token and self.hf_repo_id:\n            if filename == 'best.pt' or (filename == 'last.pt' and (epoch + 1) % 5 == 0):\n                print(f\"-> Auto-pushing {filename} to Hugging Face Hub...\")\n                try:\n                    self.hf_api.upload_file(\n                        path_or_fileobj=filepath,\n                        path_in_repo=f\"./{filename}\",\n                        repo_id=self.hf_repo_id,\n                        commit_message=f\"Update {filename}: Epoch {epoch+1}\"\n                    )\n                    print(\"-> Push success!\")\n                except Exception as e:\n                    print(f\"-> Push failed: {e}\")\n\n        if filename == 'best.pt':\n             print(f\"-> Saved New Best Model! (Metric: {self.best_metric:.4f})\")\n\n    def load_checkpoint(self, filename='last.pt', load_scheduler=True, load_lr=True, force_download=False):\n        local_filepath = os.path.join(self.save_dir, filename)\n\n        if force_download or (not os.path.isfile(local_filepath) and self.hf_repo_id):\n            print(f\"-> Looking for '{filename}' on Hugging Face...\")\n            downloaded_path = None\n            try:\n                downloaded_path = hf_hub_download(\n                    repo_id=self.hf_repo_id,\n                    filename=filename,\n                    token=self.hf_token,\n                    local_dir=self.save_dir\n                )\n            except Exception:\n                try:\n                    downloaded_path = hf_hub_download(\n                        repo_id=self.hf_repo_id,\n                        filename=f\"checkpoints/{filename}\",\n                        token=self.hf_token,\n                        local_dir=self.save_dir\n                    )\n                except Exception:\n                    print(\"-> Not found on HF.\")\n\n            if downloaded_path is None:\n                print(\"-> Checkpoint not found. Starting from scratch.\")\n                return 0\n\n        target_file = local_filepath\n        if not os.path.isfile(target_file):\n             return 0\n\n        print(f\"-> Loading checkpoint from '{target_file}'...\")\n        try:\n            checkpoint = torch.load(target_file, map_location=self.device)\n            state_dict = checkpoint['state_dict']\n\n            if hasattr(self.model, \"_orig_mod\"):\n                target_model = self.model._orig_mod\n            else:\n                target_model = self.model\n            try:\n                target_model.load_state_dict(state_dict)\n                print(\">> Load state_dict success (Strict mode).\")\n            except RuntimeError as e:\n                print(f\"Warning loading state_dict: {e}\")\n                print(\">> Load with strict=False\")\n                target_model.load_state_dict(state_dict, strict=False)\n\n            if 'optimizer' in checkpoint:\n                if load_lr:\n                    self.optimizer.load_state_dict(checkpoint['optimizer'])\n                else:\n                    current_lrs = [g['lr'] for g in self.optimizer.param_groups]\n                    self.optimizer.load_state_dict(checkpoint['optimizer'])\n                    for i, g in enumerate(self.optimizer.param_groups):\n                        g['lr'] = current_lrs[i]\n\n            if 'scaler' in checkpoint: self.scaler.load_state_dict(checkpoint['scaler'])\n            if self.scheduler and checkpoint.get('scheduler') and load_scheduler:\n                self.scheduler.load_state_dict(checkpoint['scheduler'])\n\n            self.best_metric = checkpoint.get('best_metric', float('inf') if self.mode == 'min' else -float('inf'))\n            start_epoch = checkpoint['epoch'] + 1\n\n            print(f\"-> Resumed! Start Epoch: {start_epoch+1}, Best Metric: {self.best_metric:.4f}\")\n            return start_epoch\n        except Exception as e:\n            print(f\"-> Error loading checkpoint: {e}\")\n            return 0\n\n    def _check_early_stopping(self, val_metric):\n        is_improvement = False\n        if self.mode == 'min':\n            if val_metric < (self.best_metric - self.threshold):\n                is_improvement = True\n        else:\n            if val_metric > (self.best_metric + self.threshold):\n                is_improvement = True\n\n        if is_improvement:\n            self.best_metric = val_metric\n            self.patience_counter = 0\n            return False, True\n        else:\n            self.patience_counter += 1\n            print(f\"   -> EarlyStopping counter: {self.patience_counter}/{self.patience}\")\n            if self.patience_counter >= self.patience:\n                return True, False\n            return False, False\n\n    def fit(self, epochs, warmup=0, resume=False, cp_path='last.pt', force_download=False):\n        print(f\"{'='*30} START TRAINING {'='*30}\")\n\n        start_epoch = 0\n        if resume:\n            start_epoch = self.load_checkpoint(cp_path, force_download=force_download)\n\n        for epoch in range(start_epoch, epochs):\n            train_loss, train_metric, train_time = self.train_one_epoch(epoch)\n            val_loss, val_metric = self.evaluate(self.val_loader, mode=\"Val\")\n\n            curr_lr = self.optimizer.param_groups[0]['lr']\n\n            if self.scheduler and epoch >= warmup:\n                if isinstance(self.scheduler, optim.lr_scheduler.ReduceLROnPlateau):\n                    self.scheduler.step(val_metric)\n                else:\n                    self.scheduler.step()\n\n            print(f\"Epoch {epoch+1}/{epochs} | \"\n                  f\"Time: {train_time:.1f}s | \"\n                  f\"L.Train: {train_loss:.4f} | \"\n                  f\"L.Val: {val_loss:.4f} | \"\n                  f\"Met.Train: {train_metric:.2f} | \"\n                  f\"Met.Val: {val_metric:.2f} | \"\n                  f\"LR: {curr_lr:.2e}\")\n\n            \n\n            if (epoch + 1) >= self.min_epochs:\n                should_stop, is_best = self._check_early_stopping(val_metric)\n\n                if is_best:\n                    print(f\"   -> New Best Model! ({val_metric:.4f})\")\n                    self.save_checkpoint(epoch, filename='best.pt')\n\n                if should_stop:\n                    print(f\"\\n[!] Early stopping at epoch {epoch+1}\")\n                    break\n            else:\n                better = (val_metric < self.best_metric) if self.mode == 'min' else (val_metric > self.best_metric)\n                if better:\n                    self.best_metric = val_metric\n                    self.save_checkpoint(epoch, filename='best.pt')\n            self.save_checkpoint(epoch, filename='last.pt')\n\n        print(f\"\\nTraining Finished! Best Metric: {self.best_metric:.4f}\")","metadata":{"id":"WNQdYqEuqSMw","trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:39:03.109499Z","iopub.execute_input":"2026-02-03T13:39:03.109824Z","iopub.status.idle":"2026-02-03T13:39:03.145969Z","shell.execute_reply.started":"2026-02-03T13:39:03.109798Z","shell.execute_reply":"2026-02-03T13:39:03.145286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_perplexity(loss_value):\n    try:\n        return min(math.exp(loss_value), 1e5)\n    except OverflowError:\n        return float('inf')","metadata":{"id":"w5Nr04um2XxK","trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:39:10.352289Z","iopub.execute_input":"2026-02-03T13:39:10.352585Z","iopub.status.idle":"2026-02-03T13:39:10.35729Z","shell.execute_reply.started":"2026-02-03T13:39:10.352561Z","shell.execute_reply":"2026-02-03T13:39:10.356286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nhf_token = user_secrets.get_secret(\"hf_token\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:39:11.637071Z","iopub.execute_input":"2026-02-03T13:39:11.637383Z","iopub.status.idle":"2026-02-03T13:39:11.712671Z","shell.execute_reply.started":"2026-02-03T13:39:11.637355Z","shell.execute_reply":"2026-02-03T13:39:11.712002Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def count_parameters(model):\n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    \n    print(f\"Total Params: {total_params:,} ({total_params/1e6:.2f}M)\")\n    print(f\"Trainable Params: {trainable_params:,} ({trainable_params/1e6:.2f}M)\")\n    print(f\"Non-trainable Params: {total_params - trainable_params:,}\")\n    \n    return total_params, trainable_params\ncount_parameters(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:39:29.023931Z","iopub.execute_input":"2026-02-03T13:39:29.024239Z","iopub.status.idle":"2026-02-03T13:39:29.03231Z","shell.execute_reply.started":"2026-02-03T13:39:29.024213Z","shell.execute_reply":"2026-02-03T13:39:29.03147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = 10\noptimizer = torch.optim.AdamW(model.parameters(), lr=8e-5, weight_decay=5e-2)\nlr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs, eta_min=1e-5)\ncriterion = nn.CrossEntropyLoss(ignore_index=-100)\nmetric = calculate_perplexity\ntrainer = Trainer(\n    model=model,\n    train_loader=dloader_train,\n    val_loader=dloader_test,\n    optimizer=optimizer,\n    scheduler=lr_scheduler,\n    criterion=criterion,\n    metric=metric,\n    mode='min',\n    device='cuda',\n    save_dir='./checkpoints',\n    early_stopping_patience=2,\n    early_stopping_threshold=0.0,\n    early_stopping_min_epochs=5,\n    gradient_accumulation_steps=4,\n    max_grad_norm=1.0,\n    use_flash=True,\n    hf_repo_id=\"CBG6682/Glue_qqp\",\n    hf_token=hf_token\n)\ntrainer.fit(epochs, warmup=3)","metadata":{"id":"fP6zjLhU4SKU","trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:40:38.646091Z","iopub.execute_input":"2026-02-03T13:40:38.646753Z","iopub.status.idle":"2026-02-03T13:42:00.428589Z","shell.execute_reply.started":"2026-02-03T13:40:38.646727Z","shell.execute_reply":"2026-02-03T13:42:00.427458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n\n# def inspect_dataset(dataloader, tokenizer):\n#     print(\"=== KIỂM TRA DỮ LIỆU (DEEP INSPECTION) ===\")\n    \n#     # Lấy 1 batch đầu tiên\n#     for batch in dataloader:\n#         # Hỗ trợ cả trường hợp batch là dict hoặc list (tùy collate_fn)\n#         if isinstance(batch, dict):\n#             input_ids = batch[\"input_ids\"]\n#             labels = batch[\"labels\"]\n#             attention_mask = batch[\"attention_mask\"]\n#             # Kiểm tra xem có decoder_input_ids không (cho custom model)\n#             decoder_input_ids = batch.get(\"decoder_input_ids\", None)\n#         else:\n#             print(\"Batch không phải là Dictionary, vui lòng kiểm tra collate_fn\")\n#             return\n\n#         batch_size = input_ids.shape[0]\n#         print(f\"Batch Size: {batch_size}\")\n        \n#         # Chỉ in mẫu đầu tiên để soi kỹ\n#         samples_to_print = 1 \n        \n#         for i in range(samples_to_print):\n#             print(f\"\\n{'='*20} SAMPLE {i+1} {'='*20}\")\n            \n#             # --- PHẦN 1: KIỂM TRA INPUT (ENCODER) ---\n#             print(\"\\n1. ENCODER INPUT:\")\n            \n#             # A. Decode thông thường\n#             decoded_input = tokenizer.decode(input_ids[i], skip_special_tokens=False)\n#             print(f\"   [Full Text]: {repr(decoded_input)}\")\n            \n#             # B. Soi Attention Mask (Quan trọng để xem Padding có bị che không)\n#             # Chuyển ID sang Tokens (subwords) để xem từng token\n#             raw_input_tokens = tokenizer.convert_ids_to_tokens(input_ids[i])\n#             visual_input_mask = []\n            \n#             for token, mask_val in zip(raw_input_tokens, attention_mask[i]):\n#                 if mask_val == 1:\n#                     visual_input_mask.append(token) # Giữ nguyên nếu mask = 1\n#                 else:\n#                     visual_input_mask.append(f\"[MASKED_PAD]\") # Thay thế nếu mask = 0\n            \n#             print(f\"   [Mask Check]: {' '.join(visual_input_mask)}\")\n#             print(f\"   -> (Nếu bạn thấy [MASKED_PAD] đè lên chữ thật thì là lỗi. Nó chỉ nên đè lên token pad)\")\n\n#             # --- PHẦN 2: KIỂM TRA LABELS (DECODER TARGET) ---\n#             print(\"\\n2. DECODER LABELS (TARGET):\")\n            \n#             # Xử lý -100 để decode hiển thị được\n#             lbl_ids = labels[i].clone()\n#             # Thay thế tạm -100 bằng pad_token_id để hàm decode không bị lỗi\n#             clean_lbl_ids = lbl_ids.clone()\n#             clean_lbl_ids[clean_lbl_ids == -100] = tokenizer.pad_token_id\n            \n#             # A. Decode thông thường (như con người đọc)\n#             decoded_label = tokenizer.decode(clean_lbl_ids, skip_special_tokens=False)\n#             print(f\"   [Full Text]: {repr(decoded_label)}\")\n            \n#             # B. Soi Mask -100 (Quan trọng để xem Loss có tính đúng chỗ không)\n#             # Lấy token dạng chuỗi từ cái đã clean\n#             raw_label_tokens = tokenizer.convert_ids_to_tokens(clean_lbl_ids)\n#             visual_label_mask = []\n            \n#             count_ignore = 0\n#             count_train = 0\n            \n#             for token, lbl_val in zip(raw_label_tokens, labels[i]):\n#                 if lbl_val == -100:\n#                     visual_label_mask.append(\"[IGNORE_LOSS]\")\n#                     count_ignore += 1\n#                 else:\n#                     visual_label_mask.append(token)\n#                     count_train += 1\n            \n#             print(f\"   [Loss Mask Check]: {' '.join(visual_label_mask)}\")\n#             print(f\"   -> Tokens được học: {count_train} | Tokens bị bỏ qua (-100): {count_ignore}\")\n\n#             # --- PHẦN 3: KIỂM TRA DECODER INPUT (NẾU CÓ) ---\n#             if decoder_input_ids is not None:\n#                 print(\"\\n3. DECODER INPUT (TEACHER FORCING):\")\n#                 dec_input_text = tokenizer.decode(decoder_input_ids[i], skip_special_tokens=False)\n#                 print(f\"   [Full Text]: {repr(dec_input_text)}\")\n                \n#                 # Kiểm tra nhanh: Decoder Input có phải là Labels dịch phải không?\n#                 # Token đầu tiên của Decoder Input thường là Pad hoặc Start\n#                 first_token = decoder_input_ids[i][0].item()\n#                 print(f\"   [Start Token ID]: {first_token} (Thường là {tokenizer.pad_token_id} hoặc {tokenizer.eos_token_id})\")\n\n#         break # Dừng sau batch đầu tiên\n\n# if __name__ == \"__main__\":\n#     tokenizer = AutoTokenizer.from_pretrained(\"answerdotai/ModernBERT-base\")\n\n#     if tokenizer.pad_token is None:\n#         tokenizer.pad_token = tokenizer.eos_token\n#     inspect_dataset(dloader_train, tokenizer)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:38:58.687581Z","iopub.status.idle":"2026-02-03T13:38:58.687938Z","shell.execute_reply.started":"2026-02-03T13:38:58.687758Z","shell.execute_reply":"2026-02-03T13:38:58.687774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n\n# def run_inference(model, tokenizer, input_text, device=\"cuda\"):\n#     generated_text = model.generate(\n#         tokenizer=tokenizer,\n#         text=input_text,\n#         max_new_tokens=256,   \n#         temperature=0.1, \n#         top_k=2,              \n#         repetition_penalty=1.15,\n#         device=device,\n#         verbose=True\n#     )\n#     return generated_text\n\n# test_sentences = [\n#     \"I am very happy with the results of the project.\",\n#     \"She runs very fast.\",\n#     \"It is important to drink water every day.\",\n#     \"The meeting was cancelled because the boss was sick.\",\n#     \"Although it was raining heavily, they went out for a walk.\",\n#     \"If you don't study hard, you will fail the exam.\",\n#     \"The company's revenue declined significantly due to the economic recession.\",\n#     \"Artificial intelligence is transforming the way we work and live.\",\n#     \"Please ensure that all documents are signed before submission.\",\n#     \"The test was a piece of cake.\"\n# ]\n# if __name__ == \"__main__\":\n#     tokenizer = dtrain.tokenizer\n#     model = trainer.model\n#     DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    \n#     print(f\"Using device: {DEVICE}\")\n\n#     print(\"\\n\" + \"=\"*40)\n#     text = input(\"Nhập câu gốc (hoặc 'exit' để thoát): \")\n    \n#     # if text.lower() in ['exit', 'quit']:\n#     #     break\n        \n#     # if not text.strip():\n#     #     continue\n#     result = run_inference(model, tokenizer, text, device=DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-03T13:38:58.688968Z","iopub.status.idle":"2026-02-03T13:38:58.689273Z","shell.execute_reply.started":"2026-02-03T13:38:58.689122Z","shell.execute_reply":"2026-02-03T13:38:58.689137Z"}},"outputs":[],"execution_count":null}]}