{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":87793,"databundleVersionId":11553390,"sourceType":"competition"},{"sourceId":10855324,"sourceType":"datasetVersion","datasetId":6742586},{"sourceId":11031941,"sourceType":"datasetVersion","datasetId":6869784},{"sourceId":224830487,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"🏭 GASTOS EJECUTIVOS\n\nCapex significa gasto de capital y se refiere a los costos asociados con la adquisición, el mantenimiento o la mejora de activos fijos como propiedades, planta y equipo.\n\nEstos gastos se pueden depreciar con el tiempo.\n\n➕Pros:\nLas inversiones en Capex proporcionan beneficios a largo plazo en términos de aumento de la productividad, la eficiencia y la competitividad.\n\nTambién ayudan a construir la base de activos de una empresa que puede proporcionar ingresos adicionales en el futuro.\n\n➖Contras:\nLas inversiones en Capex implican costos iniciales sustanciales que pueden afectar las finanzas de una organización si no se administran adecuadamente.\n\nAdemás, no hay garantía de que la inversión produzca los resultados deseados en términos de aumento de los beneficios o ventaja competitiva.\n\nEjemplos de Capex:\n\nEl gasto de capital incluye la compra de un nuevo edificio de fábrica, la renovación de uno existente o la compra de maquinaria.\n\n🔩Gastos operativos\n\nOpex significa gasto operativo y se refiere a los costos diarios de funcionamiento de un negocio, como salarios, impuestos, alquiler, servicios públicos, etc.\n\nA menudo se hace referencia a estos gastos como gastos \"corrientes\", ya que se incurre en ellos de forma continua.\n\n➕Pros:\nLos costes operativos pueden gestionarse más fácilmente que las inversiones en gastos de capital, ya que no implican grandes costes iniciales y no son necesarios compromisos a largo plazo.\n\nTambién proporcionan la flexibilidad para ajustar los niveles de gasto en respuesta a las condiciones cambiantes del mercado o a las necesidades de los clientes.\n\n➖Contras:\nLos costos de Opex tienden a aumentar con el tiempo debido a la inflación y otros factores que pueden afectar las finanzas de una organización si no se administran adecuadamente.\n\nAdemás, los gastos operativos pueden carecer de los beneficios a largo plazo que conllevan las inversiones en gastos de capital en términos de aumento de la productividad, la eficiencia y la competitividad.\n\nEjemplos de Opex:\n\nLos gastos operativos incluyen el pago de salarios, alquileres y facturas de servicios públicos, la compra de suministros y equipos de oficina, gastos de publicidad y marketing, costos de TI, honorarios legales, etc.","metadata":{}},{"cell_type":"markdown","source":"| Install Libraries ↑\nInstall libraries.","metadata":{}},{"cell_type":"markdown","source":"| Import Libraries ↑\nImport libraries.","metadata":{}},{"cell_type":"markdown","source":"| Configuration ↑\nThe next cell writes a YAML file with the different model parameters as well as other training and data configurations needed for the training stage.\n\nWe will use this file when instantiating the model and it is required to define the model's architecture.","metadata":{}},{"cell_type":"code","source":"%%writefile config.yaml\n\nlearning_rate: 0.001  # The learning rate for the optimizer\nbatch_size: 4  # Number of samples per batch\ntest_batch_size: 8  # Number of samples per batch\nepochs: 30  # Total training epochs\noptimizer: \"ranger\"  # Optimization algorithm\ndropout: 0.05  # Dropout regularization rate\nweight_decay: 0.0001\nk: 5\nninp: 256\nnlayers: 9\nnclass: 2\nntoken: 5  # AUGC + padding/N token\nnhead: 8\nuse_bpp: False\nuse_flip_aug: true\nbpp_file_folder: \"../../input/bpp_files/\"\ngradient_accumulation_steps: 1\nuse_triangular_attention: false\npairwise_dimension: 64\nuse_bpp: False\n\n#Data scaling\nuse_data_percentage: 1\nuse_dirty_data: true  # turn off for data scaling and data dropout experiments\n\n# Other configurations\nfold: 0\nnfolds: 6\ninput_dir: \"../../input/\"\ngpu_id: \"0\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:27.904911Z","iopub.execute_input":"2025-04-02T22:37:27.905276Z","iopub.status.idle":"2025-04-02T22:37:27.910697Z","shell.execute_reply.started":"2025-04-02T22:37:27.905249Z","shell.execute_reply":"2025-04-02T22:37:27.909772Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Utils ↑¶\nUtility functions.","metadata":{}},{"cell_type":"code","source":"class Config:\n    def __init__(self, **entries):\n        self.__dict__.update(entries)\n        self.entries=entries\n\n    def print(self):\n        print(self.entries)\n        \n\ndef default(val: Any, d: Any) -> Any:\n    \"\"\"\n    Returns `val` if it is not None, otherwise returns the default value `d`.\n    :param val: The primary value.\n    :param d: The default value to return if `val` is None.\n    :return: `val` if it is not None, otherwise `d`.\n    \"\"\"\n    return val if exists(val) else d\n\n\ndef exists(val: Any) -> bool:\n    \"\"\"\n    Checks whether a given value is not None.\n    :param val: The value to check.\n    :return: True if `val` is not None, otherwise False.\n    \"\"\"\n    return val is not None\n\n\ndef init_weights(m: torch.nn.Module) -> None:\n    \"\"\"\n    Initializes the weights of a given module if it is an instance of `torch.nn.Linear`. \n    Currently, the function does not apply any initialization but has commented-out \n    Xavier initialization methods.\n    :param m: The module to initialize, expected to be a `torch.nn.Linear` instance.\n    :return: None\n    \"\"\"\n    if m is not None and isinstance(m, nn.Linear):\n        pass\n\n\ndef load_config_from_yaml(file_path):\n    \"\"\"Load YAML file\"\"\"\n    with open(file_path, 'r') as file:\n        config = yaml.safe_load(file)\n    return Config(**config)\n\n\ndef sep():\n    print(\"—\"*100)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:33.724699Z","iopub.execute_input":"2025-04-02T22:37:33.724971Z","iopub.status.idle":"2025-04-02T22:37:33.730969Z","shell.execute_reply.started":"2025-04-02T22:37:33.724931Z","shell.execute_reply":"2025-04-02T22:37:33.730058Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Dropout ↑¶\nDropout functions.","metadata":{}},{"cell_type":"code","source":"class Dropout(nn.Module):\n    \"\"\"\n    Implementation of dropout with the ability to share the dropout mask\n    along a particular dimension.\n\n    If not in training mode, this module computes the identity function.\n    \"\"\"\n\n    def __init__(self, r: float, batch_dim: Union[int, List[int]]):\n        \"\"\"\n        Args:\n            r:\n                Dropout rate\n            batch_dim:\n                Dimension(s) along which the dropout mask is shared\n        \"\"\"\n        super(Dropout, self).__init__()\n\n        self.r = r\n        if type(batch_dim) == int:\n            batch_dim = [batch_dim]\n        self.batch_dim = batch_dim\n        self.dropout = nn.Dropout(self.r)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            x:\n                Tensor to which dropout is applied. Can have any shape\n                compatible with self.batch_dim\n        \"\"\"\n        shape = list(x.shape)\n        if self.batch_dim is not None:\n            for bd in self.batch_dim:\n                shape[bd] = 1\n        mask = x.new_ones(shape)\n        mask = self.dropout(mask)\n        x = x * mask\n        return x\n\n\nclass DropoutRowwise(Dropout):\n    \"\"\"\n    Convenience class for rowwise dropout as described in subsection\n    1.11.6.\n    \"\"\"\n\n    __init__ = partialmethod(Dropout.__init__, batch_dim=-3)\n\n\nclass DropoutColumnwise(Dropout):\n    \"\"\"\n    Convenience class for columnwise dropout as described in subsection\n    1.11.6.\n    \"\"\"\n\n    __init__ = partialmethod(Dropout.__init__, batch_dim=-2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:37.098694Z","iopub.execute_input":"2025-04-02T22:37:37.099015Z","iopub.status.idle":"2025-04-02T22:37:37.105445Z","shell.execute_reply.started":"2025-04-02T22:37:37.098984Z","shell.execute_reply":"2025-04-02T22:37:37.104674Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Mish ↑\nMish is a self-regularized non-monotonic activation function proposed by Diganta Misra in 2019. Its purpose is to serve as an alternative to more common activation functions like ReLU, Leaky ReLU, or Swish/SiLU in neural networks.\n\nMish has several properties that can be beneficial in neural networks:\n\nIt's smooth, unlike ReLU which has a non-differentiable point at 0\nIt's non-monotonic, which allows for better gradient flow in some contexts\nIt has a slight regularization effect due to its bounded nature at large negative values\nIt often provides better performance on various deep learning tasks compared to ReLU","metadata":{}},{"cell_type":"code","source":"class Mish(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n    def forward(self, x):\n        return x * (torch.tanh(F.softplus(x)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:41.669696Z","iopub.execute_input":"2025-04-02T22:37:41.670014Z","iopub.status.idle":"2025-04-02T22:37:41.674608Z","shell.execute_reply.started":"2025-04-02T22:37:41.669987Z","shell.execute_reply":"2025-04-02T22:37:41.673645Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| GeM Pooling ↑\nThe GeM layer (Generalized Mean Pooling) is a learnable pooling operation that allows the model to adaptively adjust the pooling behavior based on the value of the hyperparameter p. It can be used as a replacement for traditional pooling layers (such as average pooling or max pooling) in neural network architectures. GeM pooling is a generalization of average and max pooling used in deep learning. It computes the p-norm of each feature map, which makes it very useful in tasks like image retrieval and recognition.\n\nThe parameter p controls the pooling behavior:\n\nWhen p = 1: equivalent to average pooling\nAs p → ∞: approaches max pooling\nThese implementations allow p to be a learnable parameter, which lets the network determine the optimal pooling strategy during training. The layer is initialized with default values for p and eps, but these can be modified when creating an instance\n\nHere's a breakdown of the class:\n\nInitialization (__init__ method):\n\np: The parameter p is a hyperparameter that determines the type of pooling. When p is set to 1, it corresponds to average pooling. When p approaches infinity, it approximates max pooling. The default value is set to 3.\neps: A small constant (eps) is added to the input tensor before performing any operations. This is to avoid division by zero when calculating the average pooling. The default value is set to 1e-6.\nforward method:\n\nx: Input tensor to be pooled.\nThe forward method performs the GeM pooling operation on the input tensor x. It first clamps the input tensor at a minimum value of eps to avoid numerical instability. Then it raises the clamped tensor to the power of p. Finally, it applies average pooling using F.avg_pool1d over the spatial dimensions of the tensor.\nThe result is the GeM-pooled tensor.","metadata":{}},{"cell_type":"code","source":"class GeM(nn.Module):\n    \"\"\"\n    1-dimensional GeM pooling.\n    \"\"\"\n    def __init__(self, p=3, eps=1e-6):\n        super().__init__()\n        self.p = Parameter(torch.ones(1)*p)\n        self.eps = eps\n\n    def forward(self, x):\n        kernel_size = (x.size(-1))\n        output = F.avg_pool1d(\n            x.clamp(min=self.eps).pow(self.p), \n            kernel_size\n        ).pow(1./self.p)\n        return output\n\n    def __repr__(self):\n        return f'GeM(p={self.p}, eps={self.eps})'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:45.485457Z","iopub.execute_input":"2025-04-02T22:37:45.485807Z","iopub.status.idle":"2025-04-02T22:37:45.491463Z","shell.execute_reply.started":"2025-04-02T22:37:45.48577Z","shell.execute_reply":"2025-04-02T22:37:45.490437Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Scaled Dot Product Attention ↑\nScaled dot product attention was proposed by Ashish Vaswani and his colleagues at Google Brain in their groundbreaking 2017 paper \"Attention Is All You Need.\" This mechanism allows transformer models to weigh the importance of different input elements when producing each output element, enabling the network to focus on relevant parts of the input sequence regardless of distance between tokens.\n\nImage","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"class ScaledDotProductAttention(nn.Module):\n    '''\n    Scaled Dot-Product Attention module, computing attention scores based on query and key similarity.\n    '''\n    \n    def __init__(self, temperature: float, attn_dropout: float = 0.1) -> None:\n        \"\"\"\n        Initializes the Scaled Dot-Product Attention module.\n        \n        :param temperature: Scaling factor for the dot product attention scores.\n        :param attn_dropout: Dropout rate applied to attention weights.\n        \"\"\"\n        super().__init__()\n        self.temperature: float = temperature\n        self.dropout: nn.Dropout = nn.Dropout(attn_dropout)\n\n    def forward(\n        self, \n        q: torch.Tensor,  # (B, nhead, L, d_k)\n        k: torch.Tensor,  # (B, nhead, L, d_k)\n        v: torch.Tensor,  # (B, nhead, L, d_v)\n        mask: torch.Tensor | None = None,  # (B, 1, L, L) or None\n        attn_mask: torch.Tensor | None = None  # (B, 1, L, L) or None\n    ) -> tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"\n        Forward pass of the Scaled Dot-Product Attention.\n        \n        :param q: Query tensor of shape (B, nhead, L, d_k), where B is batch size, nhead is the number of attention heads,\n                  L is the sequence length, and d_k is the key/query dimension.\n        :param k: Key tensor of shape (B, nhead, L, d_k).\n        :param v: Value tensor of shape (B, nhead, L, d_v), where d_v is the value dimension.\n        :param mask: Optional bias mask tensor of shape (B, 1, L, L), used for causal masking or padding.\n        :param attn_mask: Optional attention mask tensor of shape (B, 1, L, L), where -1 values indicate positions to mask.\n        :return: Tuple containing:\n            - output (torch.Tensor): The result of the attention mechanism, shape (B, nhead, L, d_v).\n            - attn (torch.Tensor): Attention weights after softmax and dropout, shape (B, nhead, L, L).\n        \"\"\"\n        \n        attn = torch.matmul(q, k.transpose(2, 3)) / self.temperature  # (B, nhead, L, L)\n        \n        if mask is not None:\n            attn = attn + mask  # Apply bias mask (B, nhead, L, L)\n        \n        if attn_mask is not None:\n            attn = attn.float().masked_fill(attn_mask == -1, float('-1e-9'))  # Apply attention mask (B, nhead, L, L)\n        \n        attn = self.dropout(F.softmax(attn, dim=-1))  # (B, nhead, L, L)\n        output = torch.matmul(attn, v)  # (B, nhead, L, d_v)\n        \n        return output, attn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:48.887079Z","iopub.execute_input":"2025-04-02T22:37:48.887432Z","iopub.status.idle":"2025-04-02T22:37:48.893972Z","shell.execute_reply.started":"2025-04-02T22:37:48.887403Z","shell.execute_reply":"2025-04-02T22:37:48.893182Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| MultiHead Attention ↑\nThe same process as before can be repeated many times with different Key, Query, and Value projections, forming what is called a multi-head attention. Each head can focus on different projections of the input embeddings. Multihead attention extends self-attention by applying multiple attention mechanisms (or \"heads\") in parallel. Each head learns different attention patterns, which are then combined to produce a more expressive representation.\n\nImage\nInput shapes \nThe inputs q, k, and v (query, key, value) have the following shapes:\n\nq: [bs, len_q, d_model]\nk: [bs, len_k, d_model]\nv: [bs, len_v, d_model]\nWhere:\n\nbs is the batch size (first dimension)\nlen_q is the sequence length of the query\nlen_k is the sequence length of the key\nlen_v is the sequence length of the value (typically equal to len_k)\nd_model is the model's embedding dimension\nThe module then projects these inputs into multiple heads:\n\nEach head has dimension d_k for queries and keys\nEach head has dimension d_v for values\nThere are n_head different attention heads\nThe attention calculations happen in a shape of [bs, n_head, len_q, len_k] and the output has the same dimensionality as the input query: [bs, len_q, d_model].\n\nThis is a standard multi-head attention implementation where vectors are projected into multiple subspaces, attention is calculated separately in each subspace, and then the results are concatenated and projected back to the original dimension.","metadata":{}},{"cell_type":"code","source":"class MultiHeadAttention(nn.Module):\n    \"\"\"\n    Multi-Head Attention module\n    :param d_model: The number of input features.\n    :param n_head: The number of heads to use.\n    :param d_k: The dimensionality of the keys.\n    :param d_v: The dimensionality of the values.\n    :param dropout: The dropout rate to apply to the attention weights.\n    \"\"\"\n    def __init__(\n        self,\n        d_model: int,\n        n_head: int,\n        d_k: int,\n        d_v: int,\n        dropout: float = 0.1\n    ):\n        super().__init__()\n\n        self.n_head = n_head\n        self.d_k = d_k\n        self.d_v = d_v\n\n        self.w_qs = nn.Linear(d_model, n_head * d_k, bias=False)  # (d_model) -> (n_head * d_k)\n        self.w_ks = nn.Linear(d_model, n_head * d_k, bias=False)  # (d_model) -> (n_head * d_k)\n        self.w_vs = nn.Linear(d_model, n_head * d_v, bias=False)  # (d_model) -> (n_head * d_v)\n        self.fc = nn.Linear(n_head * d_v, d_model, bias=False)  # (n_head * d_v) -> (d_model)\n\n        self.attention = ScaledDotProductAttention(temperature=d_k ** 0.5)\n\n        self.dropout = nn.Dropout(dropout)\n        self.layer_norm = nn.LayerNorm(d_model, eps=1e-6)\n\n\n    def forward(\n        self, \n        q: torch.Tensor,  # Shape: [batch_size, len_q, d_model]\n        k: torch.Tensor,  # Shape: [batch_size, len_k, d_model]\n        v: torch.Tensor,  # Shape: [batch_size, len_v, d_model]\n        mask: Optional[torch.Tensor] = None,  # Optional attention mask\n        src_mask: Optional[torch.Tensor] = None  # Optional source mask\n               \n    ) -> Tuple[torch.Tensor, torch.Tensor]:  # Returns (output, attention)\n\n        d_k, d_v, n_head = self.d_k, self.d_v, self.n_head\n        bs, len_q, len_k, len_v = q.size(0), q.size(1), k.size(1), v.size(1)\n\n        residual = q  # Shape: [bs, len_q, d_model]\n\n        # Linear projections and reshape to multiple heads\n        q = self.w_qs(q).view(bs, len_q, n_head, d_k)  # Shape: [bs, len_q, n_head, d_k]\n        k = self.w_ks(k).view(bs, len_k, n_head, d_k)  # Shape: [bs, len_k, n_head, d_k]\n        v = self.w_vs(v).view(bs, len_v, n_head, d_v)  # Shape: [bs, len_v, n_head, d_v]\n\n        # Transpose for multi-head attention computation\n        q, k, v = (\n            q.transpose(1, 2),  # Shape: [bs, n_head, len_q, d_k]\n            k.transpose(1, 2),  # Shape: [bs, n_head, len_k, d_k]\n            v.transpose(1, 2)\n        )  # Shape: [bs, n_head, len_v, d_v]\n\n        if mask is not None:\n            mask = mask  # Shape remains unchanged\n\n        if src_mask is not None:\n            src_mask[src_mask == 0] = -1\n            src_mask = src_mask.unsqueeze(-1).float()  # Shape: [bs, len_k, 1]\n            attn_mask = torch.matmul(src_mask, src_mask.permute(0, 2, 1)).unsqueeze(1)  \n            # Shape: [bs, 1, len_k, len_k]\n            q, attn = self.attention(q, k, v, mask=mask, attn_mask=attn_mask)\n        else:\n            q, attn = self.attention(q, k, v, mask=mask)\n\n        # Reshape back to original format\n        q = q.transpose(1, 2).contiguous().view(bs, len_q, -1)  # Shape: [bs, len_q, n_head * d_v]\n        q = self.dropout(self.fc(q))  # Shape: [bs, len_q, d_model]\n        q += residual  # Shape: [bs, len_q, d_model]\n\n        q = self.layer_norm(q)  # Shape: [bs, len_q, d_model]\n\n        return q, attn  # Output: [bs, len_q, d_model], Attention: [bs, n_head, len_q, len_k]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:53.342043Z","iopub.execute_input":"2025-04-02T22:37:53.342343Z","iopub.status.idle":"2025-04-02T22:37:53.352434Z","shell.execute_reply.started":"2025-04-02T22:37:53.342319Z","shell.execute_reply":"2025-04-02T22:37:53.351624Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Positional Encoding ↑\nSince self-attention mechanisms do not have inherent order awareness, this encoding helps the model distinguish between different positions in a sequence. Positional embeddings are vectors that contain information about a position in the sequence. This adds information about the sequence even before attention is applied, and it allows attention to calculate relationships knowing the relative order.\n\nImage\nA detailed explanation of how it works can be found here, but a quick explanation is that we create a vector for each element representing its position with regard to every other element in the sequence. Positional encoding follows this formula which, in practice, we won’t really need to understand:\n\n","metadata":{}},{"cell_type":"code","source":"class PositionalEncoding(nn.Module):\n    def __init__(\n        self,\n        d_model: int,\n        dropout: float = 0.1,\n        max_len: int = 200\n    ):\n        super(PositionalEncoding, self).__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        pe = torch.zeros(max_len, d_model, dtype=torch.float32)\n        position = torch.arange(0, max_len, dtype=torch.float32).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2, dtype=torch.float32) * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        pe = pe.unsqueeze(0).transpose(0, 1)  # Shape: (max_len, 1, d_model)\n\n        self.register_buffer('pe', pe)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        :param x: tensor of shape (seq_len, batch_size, d_model)\n        :return: tensor of shape (seq_len, batch_size, d_model) with positional encodings added\n        \"\"\"\n        x = x + self.pe[:x.size(0), :]\n        return self.dropout(x)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:37:57.437707Z","iopub.execute_input":"2025-04-02T22:37:57.438004Z","iopub.status.idle":"2025-04-02T22:37:57.444151Z","shell.execute_reply.started":"2025-04-02T22:37:57.43798Z","shell.execute_reply":"2025-04-02T22:37:57.44319Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Outer Product Mean ↑\nThe OuterProductMean class was proposed in the paper Highly accurate protein structure prediction with AlphaFold. It computes pairwise interactions between elements in a sequence representation. It is designed to capture interactions between all pairs of positions in a sequence by computing their outer product. This is particularly useful in models that need to understand relationships between any two elements in a sequence, such as protein structure prediction models.\n\nThe MSA representation updates the pair representation through an element-wise outer product that is summed over the MSA sequence dimension. In contrast to previous work, this operation is applied within every block rather than once in the network, which enables the continuous communication from the evolving MSA representation to the pair representation.\n\nImage\n","metadata":{}},{"cell_type":"code","source":"class OuterProductMean(nn.Module):\n    \"\"\"\n    Outer Product Mean class.\n    :param in_dim: Dimensionality of the input sequence representations (default: 256).\n    :param dim_msa: Intermediate lower-dimensional representation (default: 32).\n    :param pairwise_dim: Final dimensionality of the pairwise output (default: 64).\n    \"\"\"\n    def __init__(\n        self,\n        in_dim: int = 256,\n        dim_msa: int = 32,\n        pairwise_dim: int = 64\n    ):\n        super(OuterProductMean, self).__init__()\n        self.proj_down1 = nn.Linear(in_dim, dim_msa)  # projects the input sequence representation into a lower dimensional space\n        self.proj_down2 = nn.Linear(dim_msa ** 2, pairwise_dim)  # projects the outer product representation (reshaped) to the final pairwise_dim.\n\n    def forward(\n        self,\n        seq_rep: torch.Tensor,  # shape: (batch_size, seq_length, in_dim)\n        pair_rep: torch.Tensor = None  # shape: (batch_size, seq_length, seq_length, pairwise_dim)\n    ):\n        seq_rep = self.proj_down1(seq_rep)  # output shape: (batch_size, seq_length, dim_msa)\n        outer_product = torch.einsum('bid,bjc -> bijcd', seq_rep, seq_rep)  # output shape: (batch_size, seq_length, seq_length, dim_msa, dim_msa)\n        outer_product = rearrange(outer_product, 'b i j c d -> b i j (c d)')  # flattens the last two dimensions: (batch_size, seq_length, seq_length, dim_msa ** 2).\n        outer_product = self.proj_down2(outer_product)  # output shape: (batch_size, seq_length, seq_length, pairwise_dim)\n\n        if pair_rep is not None:\n            outer_product = outer_product + pair_rep\n\n        return outer_product ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:38:00.54703Z","iopub.execute_input":"2025-04-02T22:38:00.547313Z","iopub.status.idle":"2025-04-02T22:38:00.552586Z","shell.execute_reply.started":"2025-04-02T22:38:00.547291Z","shell.execute_reply":"2025-04-02T22:38:00.551629Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Triangle Multiplicative Module ↑\nIn AlphaFold 2, the Triangle Multiplicative Module is a crucial component designed to capture geometric constraints inherent in protein structures. This module operates on pairwise residue representations, ensuring that the predicted distances between residues adhere to the triangle inequality principle, a fundamental property in Euclidean space.\n\nThe module updates the pair representation \n between residues \n and \n by considering their relationships with all other residues \n in the protein sequence. This is achieved through two symmetric operations:\n\nOutgoing Update: Aggregates information from all columns corresponding to residue \n and \n, effectively capturing how these residues jointly interact with others.\n\nIncoming Update: Aggregates information from all rows corresponding to residue \n and \n, focusing on how other residues jointly influence this pair.\n\nThese operations are analogous to nodes in graph theory, where the update of an edge depends on the nodes it connects and their shared neighbors. By incorporating these updates, the module enforces a form of structural consistency, ensuring that the predicted distances between residues are geometrically plausible.\n\n","metadata":{}},{"cell_type":"code","source":"def exists(val):\n    return val is not None\n\ndef default(val, d):\n    return val if val is not None else d\n\nclass TriangleMultiplicativeModule(nn.Module):\n    \"\"\"\n    This class is applied to the pairwise residue representations, ensuring that the predicted distances \n    between residues adhere to the triangle inequality principle.\n    \"\"\"\n    def __init__(\n        self,\n        *,\n        dim: int,\n        hidden_dim: Optional[int] = None,\n        mix: str = 'ingoing'\n    ):\n        super().__init__()\n        assert mix in {'ingoing', 'outgoing'}, 'mix must be either ingoing or outgoing'\n\n        hidden_dim = default(hidden_dim, dim)\n        self.norm = nn.LayerNorm(dim)\n\n        self.left_proj = nn.Linear(dim, hidden_dim)\n        self.right_proj = nn.Linear(dim, hidden_dim)\n        self.left_gate = nn.Linear(dim, hidden_dim)\n        self.right_gate = nn.Linear(dim, hidden_dim)\n        self.out_gate = nn.Linear(dim, hidden_dim)\n\n        # Initialize all gating to identity\n        for gate in (self.left_gate, self.right_gate, self.out_gate):\n            nn.init.constant_(gate.weight, 0.)\n            nn.init.constant_(gate.bias, 1.)\n\n        if mix == 'outgoing':\n            self.mix_einsum_eq = '... i k d, ... j k d -> ... i j d'\n        elif mix == 'ingoing':\n            self.mix_einsum_eq = '... k j d, ... k i d -> ... i j d'\n\n        self.to_out_norm = nn.LayerNorm(hidden_dim)\n        self.to_out = nn.Linear(hidden_dim, dim)\n\n    def forward(\n        self,\n        x: torch.Tensor,                  # (batch_size, seq_len, seq_len, dim)\n        src_mask: Optional[torch.Tensor] = None  # (batch_size, seq_len)\n    ) -> torch.Tensor:                    # Output: (batch_size, seq_len, seq_len, dim)\n        if exists(src_mask):\n            src_mask = src_mask.unsqueeze(-1).float()  # (batch_size, seq_len, 1)\n            mask = torch.matmul(src_mask, src_mask.permute(0, 2, 1))  # (batch_size, seq_len, seq_len)\n            mask = rearrange(mask, 'b i j -> b i j ()')  # (batch_size, seq_len, seq_len, 1)\n\n        assert x.shape[1] == x.shape[2], 'feature map must be symmetrical'\n        \n        x = self.norm(x)  # (batch_size, seq_len, seq_len, dim)\n\n        left = self.left_proj(x)  # (batch_size, seq_len, seq_len, hidden_dim)\n        right = self.right_proj(x) # (batch_size, seq_len, seq_len, hidden_dim)\n\n        if exists(src_mask):\n            left = left * mask  # (batch_size, seq_len, seq_len, hidden_dim)\n            right = right * mask # (batch_size, seq_len, seq_len, hidden_dim)\n\n        left_gate = self.left_gate(x).sigmoid()   # (batch_size, seq_len, seq_len, hidden_dim)\n        right_gate = self.right_gate(x).sigmoid() # (batch_size, seq_len, seq_len, hidden_dim)\n        out_gate = self.out_gate(x).sigmoid()     # (batch_size, seq_len, seq_len, hidden_dim)\n\n        left = left * left_gate   # (batch_size, seq_len, seq_len, hidden_dim)\n        right = right * right_gate # (batch_size, seq_len, seq_len, hidden_dim)\n\n        out = einsum(self.mix_einsum_eq, left, right)  # (batch_size, seq_len, seq_len, hidden_dim)\n\n        out = self.to_out_norm(out)  # (batch_size, seq_len, seq_len, hidden_dim)\n        out = out * out_gate         # (batch_size, seq_len, seq_len, hidden_dim)\n        return self.to_out(out)      # (batch_size, seq_len, seq_len, dim)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:38:04.777982Z","iopub.execute_input":"2025-04-02T22:38:04.778283Z","iopub.status.idle":"2025-04-02T22:38:04.787108Z","shell.execute_reply.started":"2025-04-02T22:38:04.778257Z","shell.execute_reply":"2025-04-02T22:38:04.786339Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Triangle Attention ↑\n","metadata":{}},{"cell_type":"code","source":"class TriangleAttention(nn.Module):\n    def __init__(\n        self,\n        in_dim: int = 128,\n        dim: int = 32,\n        n_heads: int = 4,\n        wise: Literal['row', 'col'] = 'row'\n    ):\n        \"\"\"\n        Implements Triangle Attention Mechanism.\n        :param in_dim: Input feature dimension.\n        :param dim: Dimension of query, key, and value per head.\n        :param n_heads: Number of attention heads.\n        :param wise: Whether to apply row-wise or column-wise attention.\n        \"\"\"\n        super(TriangleAttention, self).__init__()\n        self.n_heads = n_heads\n        self.wise = wise\n        self.norm = nn.LayerNorm(in_dim)\n        self.to_qkv = nn.Linear(in_dim, dim * 3 * n_heads, bias=False)\n        self.linear_for_pair = nn.Linear(in_dim, n_heads, bias=False)\n        self.to_gate = nn.Sequential(\n            nn.Linear(in_dim, in_dim),\n            nn.Sigmoid()\n        )\n        self.to_out = nn.Linear(n_heads * dim, in_dim)\n\n    def forward(self, z: torch.Tensor, src_mask: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Forward pass for TriangleAttention.\n        :param z: Input tensor of shape (B, I, J, in_dim).\n        :param src_mask: Source mask of shape (B, I, J).\n        :return: Output tensor of shape (B, I, J, in_dim).\n        \"\"\"\n        # Spawn pair mask\n        src_mask = src_mask.clone()\n        src_mask[src_mask == 0] = -1\n        src_mask = src_mask.unsqueeze(-1).float()  # (B, I, J, 1)\n        attn_mask = torch.matmul(src_mask, src_mask.permute(0, 2, 1))  # (B, I, J, I)\n\n        wise = self.wise\n        z = self.norm(z)  # (B, I, J, in_dim)\n\n        # Compute bias and gate\n        gate = self.to_gate(z)  # [1] (B, I, J, in_dim)\n        b = self.linear_for_pair(z)  # [5] (B, I, J, n_heads) \n\n        # Compute Q, K, V\n        q, k, v = torch.chunk(self.to_qkv(z), 3, -1)  # [2], [3], [4]: each (B, I, J, n_heads * dim)\n        q, k, v = map(lambda x: rearrange(x, 'b i j (h d)->b i j h d', h=self.n_heads), (q, k, v))  \n        # Each: (B, I, J, n_heads, dim)\n        scale = q.size(-1) ** 0.5  # Scalar\n\n        if wise == 'row':\n            eq_attn = 'brihd,brjhd->brijh'\n            eq_multi = 'brijh,brjhd->brihd'\n            b = rearrange(b, 'b i j (r h)->b r i j h', r=1)  # (B, 1, I, J, n_heads)\n            softmax_dim = 3\n            attn_mask = rearrange(attn_mask, 'b i j->b 1 i j 1')  # (B, 1, I, J, 1)\n        elif wise == 'col':\n            eq_attn = 'bilhd,bjlhd->bijlh'\n            eq_multi = 'bijlh,bjlhd->bilhd'\n            b = rearrange(b, 'b i j (l h)->b i j l h', l=1)  # (B, I, J, 1, n_heads)\n            softmax_dim = 2\n            attn_mask = rearrange(attn_mask, 'b i j->b i j 1 1')  # (B, I, J, 1, 1)\n        else:\n            raise ValueError('wise should be col or row!')\n\n        # Compute attention logits\n        logits = (torch.einsum(eq_attn, q, k) / scale + b)  # [6], [7] (B, I, J, I, n_heads) or (B, I, J, J, n_heads)\n        logits = logits.masked_fill(attn_mask == -1, float('-1e-9'))  # Apply mask\n\n        # Compute attention weights\n        attn = logits.softmax(softmax_dim)  # [8] (B, I, J, I, n_heads) or (B, I, J, J, n_heads)\n\n        # Compute attention output\n        out = torch.einsum(eq_multi, attn, v)  # [9] (B, I, J, n_heads, dim)\n        out = gate * rearrange(out, 'b i j h d-> b i j (h d)')  # [10] (B, I, J, in_dim)\n\n        # Final projection\n        z_ = self.to_out(out)  # (B, I, J, in_dim)\n\n        return z_\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:38:07.985184Z","iopub.execute_input":"2025-04-02T22:38:07.985497Z","iopub.status.idle":"2025-04-02T22:38:07.994678Z","shell.execute_reply.started":"2025-04-02T22:38:07.98547Z","shell.execute_reply":"2025-04-02T22:38:07.993778Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"| ConvTransformer Encoder ↑\nThe Evoformer block \nThe ConvTransformerEncoderLayer is pretty similar to the EvoformerBlock in the AlphaFold2 model. Since they are similar I will discuss the EvoformerBlock first. The Evoformer module of the neural network iteratively updates MSA embedding and pair representation, essentially, detecting patterns of interaction between aminoacids. The Evoformer module consists of 48 identical blocks that take MSA embedding and pair representation on input and produce their refined versions as output. In the RibonanzaNet, the number of ConvTransformerEncoderLayer layers is different, by default it is set to 9.","metadata":{}},{"cell_type":"markdown","source":"The ConvTransformerEncoderLayer \nThe ConvTransformerEncoderLayer (or Evoformer) is composed of the building blocks we have seen before, such as MultiHeadAttention, TriangleMultiplicativeMdule and TriangleAttention. This class is similar to the Transformer Encoder from the RNAdegformer model:\n\nRNAdegformer involves the Transformer encoder, whose blocks process a one-dimensional representation of the sequence. In prior work, best predictions from RNAdegformer came from supplementing standard Transformer operations with one dimensional convolutional operations, which are effective in capturing information on sequence-local motifs, and biasing the pairwise attention matrix with terms encoding sequence distance as well as the base pair probability (BPP) matrix computed by conventional secondary structure prediction methods like EternaFold.\n\nImage","metadata":{}},{"cell_type":"code","source":"class ConvTransformerEncoderLayer(nn.Module):\n    \"\"\"\n    A Transformer Encoder Layer with convolutional enhancements and pairwise feature processing.\n    \"\"\"\n    \n    def __init__(\n        self,\n        d_model: int,\n        nhead: int,\n        dim_feedforward: int,\n        pairwise_dimension: int,\n        use_triangular_attention: bool,\n        dropout: float = 0.1,\n        k: int = 3,\n    ):\n        \"\"\"\n        :param d_model: Dimension of the input embeddings\n        :param nhead: Number of attention heads\n        :param dim_feedforward: Hidden layer size in feedforward network\n        :param pairwise_dimension: Dimension of pairwise features\n        :param use_triangular_attention: Whether to use triangular attention modules\n        :param dropout: Dropout rate\n        :param k: Kernel size for the 1D convolution\n        \"\"\"\n        super(ConvTransformerEncoderLayer, self).__init__()\n\n        # === Attention Layers ===\n        self.self_attn = MultiHeadAttention(d_model, nhead, d_model // nhead, d_model // nhead, dropout=dropout)\n\n        self.linear1 = nn.Linear(d_model, dim_feedforward)\n        self.dropout = nn.Dropout(dropout)\n        self.linear2 = nn.Linear(dim_feedforward, d_model)\n\n        # === Layer Norms ===\n        self.norm1 = nn.LayerNorm(d_model)\n        self.norm2 = nn.LayerNorm(d_model)\n        self.norm3 = nn.LayerNorm(d_model)\n        \n        # === Dropout Layers ===\n        self.dropout1 = nn.Dropout(dropout)\n        self.dropout2 = nn.Dropout(dropout)\n        self.dropout3 = nn.Dropout(dropout)\n\n        self.pairwise2heads = nn.Linear(pairwise_dimension, nhead, bias=False)\n        self.pairwise_norm = nn.LayerNorm(pairwise_dimension)\n        self.activation = nn.GELU()\n\n        self.conv = nn.Conv1d(d_model, d_model, k, padding=k // 2)\n\n        self.triangle_update_out = TriangleMultiplicativeModule(dim=pairwise_dimension, mix='outgoing')\n        self.triangle_update_in = TriangleMultiplicativeModule(dim=pairwise_dimension, mix='ingoing')\n\n        self.pair_dropout_out = DropoutRowwise(dropout)\n        self.pair_dropout_in = DropoutRowwise(dropout)\n\n        self.use_triangular_attention = use_triangular_attention\n        if self.use_triangular_attention:\n            self.triangle_attention_out = TriangleAttention(\n                in_dim=pairwise_dimension,\n                dim=pairwise_dimension // 4,\n                wise='row'\n            )\n            self.triangle_attention_in = TriangleAttention(\n                in_dim=pairwise_dimension,\n                dim=pairwise_dimension // 4,\n                wise='col'\n            )\n\n            self.pair_attention_dropout_out = DropoutRowwise(dropout)\n            self.pair_attention_dropout_in = DropoutColumnwise(dropout)\n\n        self.OuterProductMean = OuterProductMean(in_dim=d_model, pairwise_dim=pairwise_dimension)\n\n        self.pair_transition = nn.Sequential(\n            nn.LayerNorm(pairwise_dimension),\n            nn.Linear(pairwise_dimension, pairwise_dimension * 4),\n            nn.ReLU(inplace=True),\n            nn.Linear(pairwise_dimension * 4, pairwise_dimension)\n        )\n\n    def forward(\n        self,\n        src: torch.Tensor,  # Shape: (batch_size, seq_len, d_model)\n        pairwise_features: torch.Tensor,  # Shape: (batch_size, seq_len, seq_len, pairwise_dimension)\n        src_mask: torch.Tensor = None,  # Shape: (batch_size, seq_len) or None\n        return_aw: bool = False\n    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"\n        Forward pass of the ConvTransformerEncoderLayer.\n\n        :param src: Input tensor of shape (batch_size, seq_len, d_model)\n        :param pairwise_features: Pairwise feature tensor of shape (batch_size, seq_len, seq_len, pairwise_dimension)\n        :param src_mask: Optional mask tensor of shape (batch_size, seq_len)\n        :param return_aw: Whether to return attention weights\n        :return: Tuple containing processed src and pairwise_features (and optionally attention weights)\n        \"\"\"\n        src = src * src_mask.float().unsqueeze(-1)  # Shape: (batch_size, seq_len, d_model)\n        res = src  # residual\n        src = src + self.conv(src.permute(0, 2, 1)).permute(0, 2, 1)  # Shape: (batch_size, seq_len, d_model)\n        src = self.norm3(src)\n\n        pairwise_bias = self.pairwise2heads(self.pairwise_norm(pairwise_features)).permute(0, 3, 1, 2)\n        src2, attention_weights = self.self_attn(src, src, src, mask=pairwise_bias, src_mask=src_mask)  # Shape: (batch_size, seq_len, d_model)\n        \n        src = src + self.dropout1(src2)\n        src = self.norm1(src)\n        src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))  # Shape: (batch_size, seq_len, d_model)\n        src = src + self.dropout2(src2)\n        src = self.norm2(src)\n\n        pairwise_features = pairwise_features + self.OuterProductMean(src)  # Shape: (batch_size, seq_len, seq_len, pairwise_dimension)\n        pairwise_features = pairwise_features + self.pair_dropout_out(self.triangle_update_out(pairwise_features, src_mask))\n        pairwise_features = pairwise_features + self.pair_dropout_in(self.triangle_update_in(pairwise_features, src_mask))\n        \n        if self.use_triangular_attention:\n            pairwise_features = pairwise_features + self.pair_attention_dropout_out(self.triangle_attention_out(pairwise_features, src_mask))\n            pairwise_features = pairwise_features + self.pair_attention_dropout_in(self.triangle_attention_in(pairwise_features, src_mask))\n        \n        pairwise_features = pairwise_features + self.pair_transition(pairwise_features)  # Shape: (batch_size, seq_len, seq_len, pairwise_dimension)\n\n        if return_aw:\n            return src, pairwise_features, attention_weights  # Shapes: (batch_size, seq_len, d_model), (batch_size, seq_len, seq_len, pairwise_dimension), (batch_size, nhead, seq_len, seq_len)\n        else:\n            return src, pairwise_features  # Shapes: (batch_size, seq_len, d_model), (batch_size, seq_len, seq_len, pairwise_dimension)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:38:12.311875Z","iopub.execute_input":"2025-04-02T22:38:12.312179Z","iopub.status.idle":"2025-04-02T22:38:12.324406Z","shell.execute_reply.started":"2025-04-02T22:38:12.312156Z","shell.execute_reply":"2025-04-02T22:38:12.323444Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Relative Positional Encoding ↑\nThe purpose of the RelativePositionalEncoding class is to compute relative positional encodings for a sequence, which can be used in models like self-attention mechanisms.","metadata":{}},{"cell_type":"markdown","source":"**","metadata":{}},{"cell_type":"code","source":"class RelativePositionalEncoding(nn.Module):\n    \"\"\"\n    Implements relative positional encoding for sequence-based models.\n    :param dim: (int) The output embedding dimension. Default is 64.\n    \"\"\"\n    \n    def __init__(self, dim: int = 64):\n        super(RelativePositionalEncoding, self).__init__()\n        self.linear = nn.Linear(17, dim)  # (17,) -> (dim,)\n\n    def forward(self, src: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Computes the relative positional encodings for a given sequence.\n\n        :param src: Input tensor of shape (B, L, D), where:\n            - B: Batch size\n            - L: Sequence length\n            - D: Feature dimension (ignored in this module)\n        :return: Relative positional encoding of shape (L, L, dim)\n        \"\"\"\n        L = src.shape[1]  # Sequence length\n        res_id = torch.arange(L, device=src.device).unsqueeze(0)  # (1, L)\n        \n        device = res_id.device\n        bin_values = torch.arange(-8, 9, device=device)  # (17,)\n\n        d = res_id[:, :, None] - res_id[:, None, :]  # (1, L, L)\n        bdy = torch.tensor(8, device=device)\n\n        # Clipping the values within the range [-8, 8]\n        d = torch.minimum(torch.maximum(-bdy, d), bdy)  # (1, L, L)\n\n        # One-hot encoding of relative positions\n        d_onehot = (d[..., None] == bin_values).float()  # (1, L, L, 17)\n\n        assert d_onehot.sum(dim=-1).min() == 1  # Ensure proper one-hot encoding\n\n        # Linear transformation to embedding space\n        p = self.linear(d_onehot)  # (1, L, L, 17) -> (1, L, L, dim)\n\n        return p.squeeze(0)  # (L, L, dim)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:38:18.355808Z","iopub.execute_input":"2025-04-02T22:38:18.356138Z","iopub.status.idle":"2025-04-02T22:38:18.362269Z","shell.execute_reply.started":"2025-04-02T22:38:18.356112Z","shell.execute_reply":"2025-04-02T22:38:18.361493Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| RibonanzaNet ↑\nRibonanza Net:\n\nRibonanzaNet bears some similarities to top-ranking Kaggle models (from Stanford Ribonanza RNA Folding competition 2024) and RNAdegformer, because it combines 1D convolutions with Transformer encoder modules. However, the pairwise representation is updated globally, unlike BPP features used in RNAdegformer, which can be seen as a pre-computed pairwise representation. Following an embedding layer that transforms RNA bases into sequence representation, RibonanzaNet spawns a pairwise representation by computing pairwise outer products from a downsampled sequence representation. Then relative positional encodings up to 8 bases apart are added to the pairwise representation. Next, RibonanzaNet processes the sequence and pairwise representation through several layers via 1D convolution, self-attention, and triangular multiplicative updates. The combination of 1D convolution and self-attention allows the model to learn interactions between RNA bases or short segments of bases (k-mers) at any sequence distance, while leveraging information in the pairwise representation. Further, the outer product mean operation updates the pairwise representation using projected outer products, and triangular multiplicative update modules operate on the pairwise representation to update each edge with two other edges starting from/ending at the two nodes of the edge being updated. It is important to note that while the RNAdegformer and other Kaggle models that use BPP features to bias self-attention have information flowing only from BPP representation to sequence representation, in RibonanzaNet, information flows not only from the pairwise representation to sequence representation but also from sequence representation back to pairwise representation.","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass RibonanzaNet(nn.Module):\n    \"\"\"\n    A transformer-based neural network for sequence processing, incorporating convolutional transformer encoder layers,\n    outer product mean operations, and relative positional encoding.\n    \"\"\"\n    def __init__(self, config: object):\n        \"\"\"\n        Initializes the RibonanzaNet model.\n        \n        :param config: Configuration object containing model hyperparameters.\n            - ninp (int): Input embedding dimension.\n            - ntoken (int): Vocabulary size for embedding layer.\n            - nclass (int): Number of output classes.\n            - nhead (int): Number of attention heads.\n            - nlayers (int): Number of transformer encoder layers.\n            - dropout (float): Dropout probability.\n            - pairwise_dimension (int): Dimension of pairwise features.\n            - use_triangular_attention (bool): Whether to use triangular attention.\n            - use_bpp (bool): Whether to use base-pairing probability features.\n            - k (int): Kernel size for convolutions in transformer layers.\n        \"\"\"\n        super(RibonanzaNet, self).__init__()\n        self.config = config\n        nhid = config.ninp * 4\n        \n        self.transformer_encoder = []\n        print(f\"Constructing {config.nlayers} ConvTransformerEncoderLayers\")\n        for i in range(config.nlayers):\n            k = config.k if i != config.nlayers - 1 else 1\n            self.transformer_encoder.append(\n                ConvTransformerEncoderLayer(\n                    d_model=config.ninp, nhead=config.nhead,\n                    dim_feedforward=nhid,\n                    pairwise_dimension=config.pairwise_dimension,\n                    use_triangular_attention=config.use_triangular_attention,\n                    dropout=config.dropout, k=k)\n            )\n        self.transformer_encoder = nn.ModuleList(self.transformer_encoder)\n        \n        self.encoder = nn.Embedding(config.ntoken, config.ninp, padding_idx=4)\n        self.decoder = nn.Linear(config.ninp, config.nclass)\n        \n        if config.use_bpp:\n            self.mask_dense = nn.Conv2d(2, config.nhead // 4, 1)\n        else:\n            self.mask_dense = nn.Conv2d(1, config.nhead // 4, 1)\n        \n        self.OuterProductMean = OuterProductMean(in_dim=config.ninp, pairwise_dim=config.pairwise_dimension)\n        self.pos_encoder = RelativePositionalEncoding(config.pairwise_dimension)\n\n    def forward(self, src: torch.Tensor, src_mask: torch.Tensor = None, return_aw: bool = False):\n        \"\"\"\n        Forward pass of the RibonanzaNet model.\n        \n        :param src: Input tensor of shape (B, L), where B is the batch size and L is the sequence length.\n        :param src_mask: Optional mask tensor of shape (B, L, L), used for attention masking.\n        :param return_aw: Boolean flag indicating whether to return attention weights.\n        :return: Output tensor of shape (B, L, nclass) if return_aw is False, or a tuple (output, attention_weights).\n        \"\"\"\n        B, L = src.shape  # (Batch size, Sequence length)\n        src = self.encoder(src).reshape(B, L, -1)  # (B, L, ninp)\n        \n        pairwise_features = self.OuterProductMean(src)  # (B, L, L, pairwise_dimension)\n        pairwise_features = pairwise_features + self.pos_encoder(src)  # (B, L, L, pairwise_dimension)\n        \n        attention_weights = []\n        for i, layer in enumerate(self.transformer_encoder):\n            if src_mask is not None:\n                if return_aw:\n                    src, aw = layer(src, pairwise_features, src_mask, return_aw=return_aw)\n                    attention_weights.append(aw)\n                else:\n                    src, pairwise_features = layer(src, pairwise_features, src_mask, return_aw=return_aw)\n            else:\n                if return_aw:\n                    src, aw = layer(src, pairwise_features, return_aw=return_aw)\n                    attention_weights.append(aw)\n                else:\n                    src, pairwise_features = layer(src, pairwise_features, return_aw=return_aw)\n        \n        output = self.decoder(src).squeeze(-1) + pairwise_features.mean() * 0  # (B, L, nclass)\n        \n        if return_aw:\n            return output, attention_weights\n        else:\n            return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:38:23.906685Z","iopub.execute_input":"2025-04-02T22:38:23.906982Z","iopub.status.idle":"2025-04-02T22:38:23.916923Z","shell.execute_reply.started":"2025-04-02T22:38:23.906958Z","shell.execute_reply":"2025-04-02T22:38:23.91605Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Build Model ↑\nLet's build the model to check everything is OK and print the output shape.","metadata":{}},{"cell_type":"code","source":"config = load_config_from_yaml(\"config.yaml\")\nmodel = RibonanzaNet(config).cuda()\nx = torch.ones(4, 128).long().cuda()\nmask = torch.ones(4, 128).long().cuda()\nmask[:,120:] = 0\nprint(f\"Output shape: {model(x,src_mask=mask).shape}\"), sep()\nmodel","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:38:31.905785Z","iopub.execute_input":"2025-04-02T22:38:31.906092Z","iopub.status.idle":"2025-04-02T22:38:32.110407Z","shell.execute_reply.started":"2025-04-02T22:38:31.906068Z","shell.execute_reply":"2025-04-02T22:38:32.109714Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Conclusions ↑\nThat's it! We have gone through all the blocks involved in the RibonanzaNet. As we have seen, it holds many similarities with some building block of AlphaFold 2, as well as the RNAdegformer and other top Kaggle solutions from the Stanford Ribonanza RNA Folding competition.\n\nShujun has also provided both a finetuning code and an inference code for starters! Go check them out.\n\nI hope this tutorial has been useful to you to better understand the RibonanzaNet network used in this competition.\n\nPlease leave any comments below and improvements for the notebook.\n\nBest luck! 🍀","metadata":{}},{"cell_type":"code","source":"import math\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport yaml\n\n\nfrom einops import rearrange, repeat, reduce\nfrom einops.layers.torch import Rearrange\nfrom functools import partialmethod\nfrom torch import einsum\nfrom torch.nn.parameter import Parameter\nfrom typing import Any, Dict, List, Literal, Optional, Tuple, Union","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:30:29.399272Z","iopub.execute_input":"2025-04-02T22:30:29.399578Z","iopub.status.idle":"2025-04-02T22:30:29.403991Z","shell.execute_reply.started":"2025-04-02T22:30:29.399548Z","shell.execute_reply":"2025-04-02T22:30:29.403196Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"we are using commandline \"python inference.py\" to verify correct installtion in this demo. You can instead use rhofold+ net api and write your python code for faster processing.\n\n1. setup rhofold+\nplease follow instruction at https://github.com/ml4bio/RhoFold\n   \n2. setup usalign\nplease follow instruction at https://github.com/pylelab/USalign  \nuse:  g++ -static -O3 -ffast-math -lm -o USalign USalign.cpp  ","metadata":{}},{"cell_type":"code","source":"try:\n    import Bio\nexcept:\n    #for rhofold+ #####################\n    !pip install biopython\n    !pip install ml-collections\n    !pip install python-box\n    !pip install dm-tree\n    !pip install openmm[cuda12]\n\n\n\n\n\nfrom copy import deepcopy\n\nimport pandas as pd\nfrom Bio.PDB import Atom, Model, Chain, Residue, Structure, PDBParser\nfrom Bio import SeqIO\nimport os, sys\nimport re\nimport numpy as np\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n\nprint('IMPORT OK !!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:30:29.404729Z","iopub.execute_input":"2025-04-02T22:30:29.404988Z","iopub.status.idle":"2025-04-02T22:30:29.419301Z","shell.execute_reply.started":"2025-04-02T22:30:29.404956Z","shell.execute_reply":"2025-04-02T22:30:29.418483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PYTHON = sys.executable\nprint('PYTHON',PYTHON)\n\nRHONET_DIR=\\\n'/kaggle/input/data-for-demo-for-rhofold-plus-with-kaggle-msa/RhoFold-main'\n#'<your downloaded rhofold repo>/RhoFold-main'\n\nUSALIGN = \\\n'/kaggle/working//USalign'\n#'<your us align path>/USalign'\n\nos.system('cp /kaggle/input/usalign/USalign /kaggle/working/')\nos.system('sudo chmod u+x /kaggle/working//USalign')\n\n\nDATA_KAGGLE_DIR = '/kaggle/input/stanford-rna-3d-folding'\nSEQ_DF = pd.read_csv(f'{DATA_KAGGLE_DIR}/train_sequences.csv')\nLABEL_DF = pd.read_csv(f'{DATA_KAGGLE_DIR}/train_labels.csv')\nLABEL_DF['target_id'] = LABEL_DF['ID'].apply(lambda x: '_'.join(x.split('_')[:-1]))\n\n\n# helper ----\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\n# visualisation helper ----\ndef set_aspect_equal(ax):\n\tx_limits = ax.get_xlim()\n\ty_limits = ax.get_ylim()\n\tz_limits = ax.get_zlim()\n\n\t# Compute the mean of each axis\n\tx_middle = np.mean(x_limits)\n\ty_middle = np.mean(y_limits)\n\tz_middle = np.mean(z_limits)\n\n\t# Compute the max range across all axes\n\tmax_range = max(x_limits[1] - x_limits[0],\n\t\t\t\t\ty_limits[1] - y_limits[0],\n\t\t\t\t\tz_limits[1] - z_limits[0]) / 2.0\n\n\t# Set the new limits to ensure equal scaling\n\tax.set_xlim(x_middle - max_range, x_middle + max_range)\n\tax.set_ylim(y_middle - max_range, y_middle + max_range)\n\tax.set_zlim(z_middle - max_range, z_middle + max_range)\n\n\n\n\n# xyz df helper --------------------\ndef get_truth_df(target_id):\n    truth_df = LABEL_DF[LABEL_DF['target_id'] == target_id]\n    truth_df = truth_df.reset_index(drop=True)\n    return truth_df\n\ndef parse_pdb_to_df(pdb_file, target_id):\n    parser = PDBParser()\n    structure = parser.get_structure('', pdb_file)\n\n    df = []\n    for model in structure:\n        for chain in model:\n            print(chain)\n            chain_data = []\n            for residue in chain:\n                # print(residue)\n                if residue.get_resname() in ['A', 'U', 'G', 'C']:\n                    # Check if the residue has a C1' atom\n                    if 'C1\\'' in residue:\n                        atom = residue['C1\\'']\n                        xyz = atom.get_coord()\n                        resname = residue.get_resname()\n                        resid = residue.get_id()[1]\n\n                        #todo detect discontinous: resid = prev_resid+1\n                        #ID\tresname\tresid\tx_1\ty_1\tz_1\n                        chain_data.append(dict(\n                            ID = target_id+'_'+str(resid),\n                            resname=resname,\n                            resid=resid,\n                            x_1=xyz[0],\n                            y_1=xyz[1],\n                            z_1=xyz[2],\n                        ))\n                        ##print(f\"Residue {resname} {resid}, Atom: {atom.get_name()}, xyz: {xyz}\")\n\n            if len(chain_data)!=0:\n                chain_df = pd.DataFrame(chain_data)\n                df.append(chain_df)\n                ##print(chain_df)\n    return df\n\n# usalign helper --------------------\ndef write_target_line(\n    atom_name, atom_serial, residue_name, chain_id, residue_num, x_coord, y_coord, z_coord, occupancy=1.0, b_factor=0.0, atom_type='P'\n):\n    \"\"\"\n    Writes a single line of PDB format based on provided atom information.\n\n    Args:\n        atom_name (str): Name of the atom (e.g., \"N\", \"CA\").\n        atom_serial (int): Atom serial number.\n        residue_name (str): Residue name (e.g., \"ALA\").\n        chain_id (str): Chain identifier.\n        residue_num (int): Residue number.\n        x_coord (float): X coordinate.\n        y_coord (float): Y coordinate.\n        z_coord (float): Z coordinate.\n        occupancy (float, optional): Occupancy value (default: 1.0).\n        b_factor (float, optional): B-factor value (default: 0.0).\n\n    Returns:\n        str: A single line of PDB string.\n    \"\"\"\n    return f'ATOM  {atom_serial:>5d}  {atom_name:<5s} {residue_name:<3s} {residue_num:>3d}    {x_coord:>8.3f}{y_coord:>8.3f}{z_coord:>8.3f}{occupancy:>6.2f}{b_factor:>6.2f}           {atom_type}\\n'\n\ndef write_xyz_to_pdb(df, pdb_file, xyz_id = 1):\n    resolved_cnt = 0\n    with open(pdb_file, 'w') as target_file:\n        for _, row in df.iterrows():\n            x_coord = row[f'x_{xyz_id}']\n            y_coord = row[f'y_{xyz_id}']\n            z_coord = row[f'z_{xyz_id}']\n\n            if x_coord > -1e17 and y_coord > -1e17 and z_coord > -1e17:\n                resolved_cnt += 1\n                target_line = write_target_line(\n                    atom_name=\"C1'\",\n                    atom_serial=int(row['resid']),\n                    residue_name=row['resname'],\n                    chain_id='0',\n                    residue_num=int(row['resid']),\n                    x_coord=x_coord,\n                    y_coord=y_coord,\n                    z_coord=z_coord,\n                    atom_type='C',\n                )\n                target_file.write(target_line)\n    return resolved_cnt\n\ndef parse_usalign_for_tm_score(output):\n    # Extract TM-score based on length of reference structure (second)\n    tm_score_match = re.findall(r'TM-score=\\s+([\\d.]+)', output)[1]\n    if not tm_score_match:\n        raise ValueError('No TM score found')\n    return float(tm_score_match)\n\ndef parse_usalign_for_transform(output):\n    # Locate the rotation matrix section\n    matrix_lines = []\n    found_matrix = False\n\n    for line in output.splitlines():\n        if \"The rotation matrix to rotate Structure_1 to Structure_2\" in line:\n            found_matrix = True\n        elif found_matrix and re.match(r'^\\d+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+\\s+[-\\d.]+$', line):\n            matrix_lines.append(line)\n        elif found_matrix and not line.strip():\n            break  # Stop parsing if an empty line is encountered after the matrix\n\n    # Parse the rotation matrix values\n    rotation_matrix = []\n    for line in matrix_lines:\n        parts = line.split()\n        row_values = list(map(float, parts[1:]))  # Skip the first column (index)\n        rotation_matrix.append(row_values)\n\n    return np.array(rotation_matrix)\n\ndef call_usalign(predict_df, truth_df, verbose=1):\n    truth_pdb = '~truth.pdb'\n    predict_pdb = '~predict.pdb'\n    write_xyz_to_pdb(predict_df, predict_pdb, xyz_id=1)\n    write_xyz_to_pdb(truth_df, truth_pdb, xyz_id=1)\n\n    command = f'{USALIGN} {predict_pdb} {truth_pdb} -atom \" C1\\'\" -m -'\n    output = os.popen(command).read()\n    if verbose==1:\n        print(output)\n    tm_score = parse_usalign_for_tm_score(output)\n    transform = parse_usalign_for_transform(output)\n    return tm_score, transform\n\n\n# msa helper --------------------\ndef read_msa(msa_file):\n    f = open(msa_file, 'r')\n    line = f.readlines()\n\n    msa = []\n    for i in range(0, len(line),2):\n        m = dotdict(\n            comment =line[i],\n            seqence =line[i+1],\n        )\n        assert(m.comment[0]=='>')\n        msa.append(m)\n    return msa\n\n\ndef write_msa(msa_file, msa):\n    line=[]\n    for m in msa:\n        line .append(m.comment)\n        line .append(m.seqence)\n\n    f = open(msa_file, 'wt')\n    f.writelines(line)\n    return msa\n \ndef msa_to_rhonet_file(msa_file, num_msa=5, out_dir='',target_id='xxx'):\n    msa = read_msa(msa_file)\n    msa0 = deepcopy(msa[0])\n    msa0.comment =f'>{target_id}\\n'\n    msa0 = [msa0]\n\n    a3m_file = f'{out_dir}/{target_id}.a3m'\n    fasta_file = f'{out_dir}/{target_id}.fasta'\n    os.makedirs(out_dir, exist_ok=True)\n\n    write_msa(fasta_file, msa0)\n    write_msa(a3m_file, msa[:num_msa])\n\nprint('HELPER OK!!!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:30:29.420159Z","iopub.execute_input":"2025-04-02T22:30:29.420436Z","iopub.status.idle":"2025-04-02T22:30:29.725741Z","shell.execute_reply.started":"2025-04-02T22:30:29.420405Z","shell.execute_reply":"2025-04-02T22:30:29.72484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#start here!!!\n\n\nout_dir   ='/kaggle/working/'\ntarget_id = '1EIY_C'.upper()\nsequence  = 'GCCGAGGUAGCUCAGUUGGUAGAGCAUGCGACUGAAAAUCGCAGUGUCCGCGGUUCGAUUCCGCGCCUCGGCACCA'\nprint('len(sequence):',len(sequence))\n\n\n#1. prepare input\nmsa_file = f'{DATA_KAGGLE_DIR}/MSA/1EIY_C.MSA.fasta'\nmsa_to_rhonet_file(msa_file, num_msa=5, out_dir=out_dir,target_id=target_id)\n\ncmd1 = f'cd {RHONET_DIR}'\n#cmd2 = f'export LD_PRELOAD=/usr/lib/x86_64-linux-gnu/libstdc++.so.6' #optional if you have lib error\ncmd3 = f'{PYTHON} inference.py --input_fas {out_dir}/{target_id}.fasta --input_a3m {out_dir}/{target_id}.a3m --output_dir {out_dir}/ --ckpt ./pretrained/model_20221010_params.pt'\n\n#follow rhofold repo, we use the cmdline:\n#'python inference.py --input_fas ./example/input/3owzA/3owzA.fasta --input_a3m ./example/input/3owzA/3owzA.a3m --output_dir ./example/output/3owzA/ --ckpt ./pretrained/model_20221010_params.pt'\n\n#do inference here!\n#output = os.popen(cmd1+';'+cmd2+';'+cmd3).read()\noutput = os.popen(cmd1+';'+cmd3).read()\nprint(output)\n\n#copy file from local results\n#local_result = '/kaggle/input/data-for-demo-for-rhofold-plus-with-kaggle-msa/rhofold_input_output/.'\n#!cp -a $local_result $out_dir\n\n#expected ouput from rtx A6000 (non ada)\n'''\n2025-03-14 20:40:50,182 - INFO: Constructing RhoFold\n2025-03-14 20:40:51,221 - INFO:     loading ./pretrained/model_20221010_params.pt\n2025-03-14 20:40:51,743 - INFO: Input_fas /media/hp/c30d34ed-0d55-4077-82dc-b56cd13dd548/2025/kaggle/stanford-rna-3d-folding/result/rhofold_00/1EIY_C.fasta\n2025-03-14 20:40:51,743 - INFO: Input_a3m /media/hp/c30d34ed-0d55-4077-82dc-b56cd13dd548/2025/kaggle/stanford-rna-3d-folding/result/rhofold_00/1EIY_C.a3m\n2025-03-14 20:40:51,743 - INFO: Started RhoFold Inference\n2025-03-14 20:40:51,755 - INFO:     Inference using device cuda\n2025-03-14 20:40:54,523 - INFO:     Export PDB file to /media/hp/c30d34ed-0d55-4077-82dc-b56cd13dd548/2025/kaggle/stanford-rna-3d-folding/result/rhofold_00//unrelaxed_model.pdb\n2025-03-14 20:40:54,523 - INFO: Finished RhoFold Inference in 2.780 seconds\n2025-03-14 20:40:54,523 - INFO: Started Amber Relaxation : 1000 iterations\n2025-03-14 20:40:54,523 - INFO:     AmberRelaxation: Using OpenCL\n2025-03-14 20:41:09,410 - INFO:     Minimizing ...\n2025-03-14 20:42:38,203 - INFO:     Energy at Minima is -505932.780 kcal/mol\n2025-03-14 20:42:38,362 - INFO:     Export PDB file to /media/hp/c30d34ed-0d55-4077-82dc-b56cd13dd548/2025/kaggle/stanford-rna-3d-folding/result/rhofold_00//relaxed_1000_model.pdb\n2025-03-14 20:42:38,363 - INFO: Finished Amber Relaxation : 1000 iterations in 103.840 seconds\n'''\n\n#expected ouput from P100\n'''\nlen(sequence): 76\n2025-03-14 16:03:33,110 - INFO: Constructing RhoFold\n2025-03-14 16:03:34,429 - INFO:     loading ./pretrained/model_20221010_params.pt\n2025-03-14 16:03:35,088 - INFO: Input_fas /kaggle/working//1EIY_C.fasta\n2025-03-14 16:03:35,089 - INFO: Input_a3m /kaggle/working//1EIY_C.a3m\n2025-03-14 16:03:35,089 - INFO: Started RhoFold Inference\n2025-03-14 16:03:35,093 - INFO:     Inference using device cuda\n2025-03-14 16:03:40,177 - INFO:     Export PDB file to /kaggle/working///unrelaxed_model.pdb\n2025-03-14 16:03:40,177 - INFO: Finished RhoFold Inference in 5.088 seconds\n2025-03-14 16:03:40,177 - INFO: Started Amber Relaxation : 1000 iterations\n2025-03-14 16:03:40,177 - INFO:     AmberRelaxation: Using OpenCL\n2025-03-14 16:04:03,056 - INFO:     Minimizing ...\n2025-03-14 16:08:49,875 - INFO:     Energy at Minima is -497896.334 kcal/mol\n2025-03-14 16:08:50,044 - INFO:     Export PDB file to /kaggle/working///relaxed_1000_model.pdb\n2025-03-14 16:08:50,047 - INFO: Finished Amber Relaxation : 1000 iterations in 309.869 seconds\n\n'''","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:30:29.726632Z","iopub.execute_input":"2025-04-02T22:30:29.726919Z","iopub.status.idle":"2025-04-02T22:34:36.97156Z","shell.execute_reply.started":"2025-04-02T22:30:29.726896Z","shell.execute_reply":"2025-04-02T22:34:36.970696Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#visualise prediction and compute tm score\n\npredict_relax_df = parse_pdb_to_df(f'{out_dir}/relaxed_1000_model.pdb', target_id)\npredict_unrelax_df = parse_pdb_to_df(f'{out_dir}/unrelaxed_model.pdb', target_id)\n\nassert len(predict_relax_df)==1\nassert len(predict_unrelax_df)==1\npredict_relax_df = predict_relax_df[0]\npredict_unrelax_df = predict_unrelax_df[0]\n\nprint(predict_relax_df)\nprint(predict_unrelax_df)\n\ntruth_df = get_truth_df(target_id)\nprint(truth_df)\n\ntm_score_relax, transform_relax = call_usalign(predict_relax_df, truth_df, verbose=1)\ntm_score_unrelax, transform_unrelax= call_usalign(predict_unrelax_df, truth_df, verbose=0)\n\nprint('tm_score_relax', tm_score_relax)\nprint('tm_score_unrelax', tm_score_unrelax)\nprint('transform_relax\\n', transform_relax)\nprint('transform_unrelax\\n', transform_unrelax)\nzz=0\n\nif 1:\n    COLOR = ['red', 'blue', 'green', 'black', 'yellow', 'cyan', 'magenta']\n    fig = plt.figure(figsize=(10, 10))\n    ax = fig.add_subplot(111, projection='3d')\n    # ax.clear()\n\n    #unrelax\n    coord = predict_unrelax_df[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n    coord = coord@transform_unrelax[:,1:].T + transform_unrelax[:,[0]].T\n    x, y, z = coord[:, 0], coord[:, 1], coord[:, 2]\n    ax.scatter(x, y, z, c='red', s=30, alpha=1)\n    ax.plot(x, y, z, color='red', linewidth=1, alpha=1, label=f'unrelax (tm:{tm_score_unrelax:0.3f})')\n\n\n    #relax\n    coord = predict_relax_df[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n    coord = coord@transform_relax[:,1:].T + transform_relax[:,[0]].T\n    x, y, z = coord[:, 0], coord[:, 1], coord[:, 2]\n    ax.scatter(x, y, z, c='orange', s=30, alpha=1)\n    ax.plot(x, y, z, color='orange', linewidth=1, alpha=1, label=f'relax (tm:{tm_score_relax:0.3f})')\n\n    # truth\n    truth = truth_df[['x_1', 'y_1', 'z_1']].to_numpy().astype('float32')\n    x, y, z = truth[:, 0], truth[:, 1], truth[:, 2]\n    ax.scatter(x, y, z, c='black', s=30, alpha=1)\n    ax.plot(x, y, z, color='black', linewidth=1, alpha=1, label=f'truth')\n\n    set_aspect_equal(ax)\n    plt.legend()\n    plt.show()\n    # plt.waitforbuttonpress()\n    plt.close()\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-02T22:34:36.972489Z","iopub.execute_input":"2025-04-02T22:34:36.972797Z","iopub.status.idle":"2025-04-02T22:34:37.305533Z","shell.execute_reply.started":"2025-04-02T22:34:36.972767Z","shell.execute_reply":"2025-04-02T22:34:37.304703Z"}},"outputs":[],"execution_count":null}]}