# Model for ONNX conversion

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Dict, Optional
import math

from model import (
    ModelArgs,
    RMSNorm,
    ReZero,
    FeedForward,
)


def precompute_freqs_cos_sins(dim: int, end: int, theta: float = 10000.0):
    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
    t = torch.arange(end, device=freqs.device)
    freqs = torch.outer(t, freqs).float()
    freqs_cos = torch.cos(freqs)  # real
    freqs_sin = torch.sin(freqs)  # imag
    freqs_cos.requires_grad = False
    freqs_sin.requires_grad = False
    return freqs_cos, freqs_sin


def apply_rotary_emb_real(
    x: torch.Tensor,
    freqs_cos: torch.Tensor,
    freqs_sin: torch.Tensor,
) -> torch.Tensor:
    seqlen, n_heads, _ = x.shape
    x_r, x_i = x.float().reshape(seqlen, n_heads, -1, 2).unbind(-1)
    _, _, head_dim = x_r.shape
    freqs_cos_v = freqs_cos.view(seqlen, 1, head_dim)
    freqs_sin_v = freqs_sin.view(seqlen, 1, head_dim)
    x_out_r = x_r * freqs_cos_v - x_i * freqs_sin_v
    x_out_i = x_r * freqs_sin_v + x_i * freqs_cos_v
    x_out = torch.stack((x_out_r, x_out_i), dim=-1).flatten(2)
    return x_out.type_as(x)


def apply_rotary_emb_real_batched(
    x: torch.Tensor,
    freqs_cos: torch.Tensor,
    freqs_sin: torch.Tensor,
) -> torch.Tensor:
    bsz, seqlen, n_heads, _ = x.shape
    x_r, x_i = x.float().reshape(bsz, seqlen, n_heads, -1, 2).unbind(-1)
    _, _, _, head_dim = x_r.shape
    freqs_cos_v = freqs_cos.view(1, seqlen, 1, head_dim)
    freqs_sin_v = freqs_sin.view(1, seqlen, 1, head_dim)
    x_out_r = x_r * freqs_cos_v - x_i * freqs_sin_v
    x_out_i = x_r * freqs_sin_v + x_i * freqs_cos_v
    x_out = torch.stack((x_out_r, x_out_i), dim=-1).flatten(3)
    return x_out.type_as(x)


class Attention(nn.Module):
    def __init__(self, args: ModelArgs):
        super().__init__()

        self.n_heads = args.n_heads
        self.head_dim = args.dim // args.n_heads

        self.wq = nn.Linear(
            args.dim,
            self.n_heads * self.head_dim,
            bias=False,
            dtype=torch.float32,
        )
        self.wk = nn.Linear(
            args.dim,
            self.n_heads * self.head_dim,
            bias=False,
            dtype=torch.float32,
        )
        self.wv = nn.Linear(
            args.dim,
            self.n_heads * self.head_dim,
            bias=False,
            dtype=torch.float32,
        )
        self.wo = nn.Linear(
            self.n_heads * self.head_dim,
            args.dim,
            bias=False,
            dtype=torch.float32,
        )

        self.cache = args.cache

    def forward(
        self,
        x: torch.Tensor,
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        mask: torch.Tensor,
    ):
        bsz, seqlen, _ = x.shape
        xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)

        xq = xq.view(bsz, seqlen, self.n_heads, self.head_dim)
        xk = xk.view(bsz, seqlen, self.n_heads, self.head_dim)
        xv = xv.view(bsz, seqlen, self.n_heads, self.head_dim)

        xq = apply_rotary_emb_real_batched(
            xq,
            freqs_cos=freqs_cos[:seqlen],
            freqs_sin=freqs_sin[:seqlen],
        )
        xk = apply_rotary_emb_real_batched(
            xk,
            freqs_cos=freqs_cos[:seqlen],
            freqs_sin=freqs_sin[:seqlen],
        )

        keys = xk
        values = xv

        xq = xq.transpose(1, 2)
        keys = torch.permute(keys, (0, 2, 3, 1))
        values = values.transpose(1, 2)
        scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim)
        scores = scores + mask
        scores = F.softmax(scores.float(), dim=-1).type_as(xq)
        output = torch.matmul(scores, values)  # (bsz, n_local_heads, slen, head_dim)
        output = output.transpose(1, 2)

        return self.wo(output.contiguous().view(bsz, seqlen, -1)), xk, xv


class AttentionOneStep(Attention):
    def __init__(self, args: ModelArgs):
        super().__init__(args)

    def forward(
        self,
        x: torch.Tensor,
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        xk_in: torch.Tensor,
        xv_in: torch.Tensor,
        mask: torch.Tensor,
    ):
        seqlen, bsz, _, _ = xk_in.shape
        _, new_seqlen, _ = x.shape

        xk = xk_in.transpose(0, 1)
        xv = xv_in.transpose(0, 1)

        xq = self.wq(x).view(bsz, new_seqlen, self.n_heads, self.head_dim)
        xk_new = self.wk(x).view(bsz, new_seqlen, self.n_heads, self.head_dim)
        xv_new = self.wv(x).view(bsz, new_seqlen, self.n_heads, self.head_dim)
        xq = apply_rotary_emb_real_batched(
            xq,
            freqs_cos=freqs_cos[seqlen : seqlen + new_seqlen],
            freqs_sin=freqs_sin[seqlen : seqlen + new_seqlen],
        )
        xk_new = apply_rotary_emb_real_batched(
            xk_new,
            freqs_cos=freqs_cos[seqlen : seqlen + new_seqlen],
            freqs_sin=freqs_sin[seqlen : seqlen + new_seqlen],
        )

        xq = xq.transpose(1, 2)
        keys = torch.cat((xk, xk_new), dim=1)
        keys = torch.permute(keys, (0, 2, 3, 1))
        values = torch.cat((xv, xv_new), dim=1).transpose(1, 2)

        scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim)
        scores = scores + mask
        scores = F.softmax(scores.float(), dim=-1).type_as(xq)
        output = torch.matmul(scores, values)  # (bsz, n_local_heads, slen, head_dim)
        output = output.transpose(1, 2)

        return (
            self.wo(output.contiguous().view(bsz, new_seqlen, -1)),
            xk_new.transpose(0, 1),
            xv_new.transpose(0, 1),
        )


class CrossAttention(nn.Module):
    def __init__(self, args: ModelArgs):
        super().__init__()

        self.n_heads = args.n_heads
        self.head_dim = args.dim // args.n_heads
        self.embedding_dim = args.cross_attention_embedding_dim

        self.wq = nn.Linear(
            args.dim,
            self.n_heads * self.head_dim,
            bias=False,
            dtype=torch.float32,
        )
        self.wk = nn.Linear(
            self.embedding_dim,
            self.n_heads * self.head_dim,
            bias=False,
            dtype=torch.float32,
        )
        self.wv = nn.Linear(
            self.embedding_dim,
            self.n_heads * self.head_dim,
            bias=False,
            dtype=torch.float32,
        )
        self.wo = nn.Linear(
            self.n_heads * self.head_dim,
            args.dim,
            bias=False,
            dtype=torch.float32,
        )

    def forward(
        self,
        x: torch.Tensor,
        encoder_out: torch.Tensor,
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        mask: torch.Tensor,
    ):
        bsz, seqlen, dim = x.shape
        ebsz, encoder_seqlen, edim = encoder_out.shape
        assert edim == self.embedding_dim
        assert bsz == ebsz

        xk, xv = self.wk(encoder_out), self.wv(encoder_out)
        xk = xk.view(ebsz, encoder_seqlen, self.n_heads, self.head_dim)
        xv = xv.view(ebsz, encoder_seqlen, self.n_heads, self.head_dim)
        xk = apply_rotary_emb_real_batched(
            xk,
            freqs_cos=freqs_cos[:encoder_seqlen],
            freqs_sin=freqs_sin[:encoder_seqlen],
        )

        xq = self.wq(x)
        xq = xq.view(bsz, seqlen, self.n_heads, self.head_dim)
        xq = apply_rotary_emb_real_batched(
            xq,
            freqs_cos=freqs_cos[:seqlen],
            freqs_sin=freqs_sin[:seqlen],
        )

        xq = xq.transpose(1, 2)
        keys = torch.permute(xk, (0, 2, 3, 1)).contiguous()
        values = xv.transpose(1, 2).contiguous()
        scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim)
        scores = (
            scores + mask
        )  # (ebsz, n_heads, seqlen, encoder_seqlen) + (ebsz, 1, seqlen, encoder_seqlen)
        scores = F.softmax(scores.float(), dim=-1).type_as(xq)
        output = torch.matmul(scores, values)  # (bsz, n_heads, seqlen, head_dim)
        output = output.transpose(1, 2)

        # keys: (ebsz, n_heads, head_dim, encoder_seqlen)
        # values: (ebsz, n_heads, encoder_seqlen, head_dim)
        return self.wo(output.contiguous().view(bsz, seqlen, -1)), keys, values


class CrossAttentionOneStep(CrossAttention):
    def __init__(self, args: ModelArgs):
        super().__init__(args)
        self.n_heads = args.n_heads
        self.head_dim = args.dim // args.n_heads

    def forward(
        self,
        x: torch.Tensor,
        start_pos: int,
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        keys: torch.Tensor,
        values: torch.Tensor,
        mask: torch.Tensor,
    ):
        bsz, new_seqlen, _ = x.shape

        xq = self.wq(x)
        xq = xq.view(bsz, new_seqlen, self.n_heads, self.head_dim)
        xq = apply_rotary_emb_real_batched(
            xq,
            freqs_cos=freqs_cos[start_pos : start_pos + new_seqlen],
            freqs_sin=freqs_sin[start_pos : start_pos + new_seqlen],
        )

        xq = xq.transpose(1, 2)
        scores = torch.matmul(xq, keys) / math.sqrt(self.head_dim)
        scores = (
            scores + mask
        )  # (ebsz, n_heads, new_seqlen, encoder_seqlen) + (ebsz, 1, new_seqlen, encoder_seqlen)
        scores = F.softmax(scores.float(), dim=-1).type_as(xq)
        output = torch.matmul(scores, values)  # (bsz, n_heads, new_seqlen, head_dim)
        output = output.transpose(1, 2)

        return self.wo(output.contiguous().view(bsz, new_seqlen, -1))


class TransformerBlock(nn.Module):
    def __init__(self, layer_id: int, args: ModelArgs):
        super().__init__()
        self.n_heads = args.n_heads
        self.dim = args.dim
        self.head_dim = args.dim // args.n_heads
        self.attention = Attention(args)
        self.feed_forward = FeedForward(
            dim=args.dim,
            hidden_dim=4 * args.dim,
            multiple_of=args.multiple_of,
            dropout_p=0.0,
        )
        self.layer_id = layer_id
        self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps)
        self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps)
        self.rezero = ReZero()
        self.enable_cross_attention = args.enable_cross_attention
        if args.enable_cross_attention:
            self.cross_attention = CrossAttention(args)
            self.cross_attention_norm = RMSNorm(args.dim, eps=args.norm_eps)
        else:
            self.cross_attention = None
            self.cross_attention_norm = None

    def forward(
        self,
        x: torch.Tensor,
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        mask: torch.Tensor,
        encoder_out: Optional[torch.Tensor],
        encoder_mask: Optional[torch.Tensor],
    ):
        a_out, a_xk, a_xv = self.attention.forward(
            self.attention_norm(x), freqs_cos, freqs_sin, mask
        )
        h = x + self.rezero(a_out)
        if self.cross_attention is not None:
            assert encoder_out is not None and encoder_mask is not None
            c_out, c_xk, c_xv = self.cross_attention.forward(
                self.cross_attention_norm(h),
                encoder_out,
                freqs_cos,
                freqs_sin,
                encoder_mask,
            )
            h = h + self.rezero(c_out)
        else:
            c_xk, c_xv = None, None
        return (
            h + self.rezero(self.feed_forward.forward(self.ffn_norm(h))),
            a_xk,
            a_xv,
            c_xk,
            c_xv,
        )


class TransformerBlockOneStep(TransformerBlock):
    def __init__(self, layer_id: int, args: ModelArgs):
        super().__init__(layer_id, args)
        self.attention = AttentionOneStep(args)
        if args.enable_cross_attention:
            self.cross_attention = CrossAttentionOneStep(args)

    def forward(
        self,
        x: torch.Tensor,
        a_xk: torch.Tensor,
        a_xv: torch.Tensor,
        c_xk: Optional[torch.Tensor],
        c_xv: Optional[torch.Tensor],
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        mask: torch.Tensor,
        encoder_mask: Optional[torch.Tensor],
    ):
        start_pos, _, _, _ = a_xk.shape
        a_out, a_xk_new, a_xv_new = self.attention.forward(
            self.attention_norm(x), freqs_cos, freqs_sin, a_xk, a_xv, mask
        )
        h = x + self.rezero(a_out)
        if self.cross_attention is not None:
            assert c_xk is not None and c_xv is not None and encoder_mask is not None
            c_out = self.cross_attention.forward(
                self.cross_attention_norm(h),
                start_pos,
                freqs_cos,
                freqs_sin,
                c_xk,
                c_xv,
                encoder_mask,
            )
            h = h + self.rezero(c_out)
        return (
            h + self.rezero(self.feed_forward.forward(self.ffn_norm(h))),
            a_xk_new,
            a_xv_new,
        )


class TransformerBlockSequence(nn.Module):
    def __init__(self, module_list):
        super().__init__()
        self.layers = module_list

    def forward(
        self,
        x: torch.Tensor,
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        mask: torch.Tensor,
        encoder_out: Optional[torch.Tensor],
        encoder_mask: Optional[torch.Tensor],
    ):
        a_xks, a_xvs, c_xks, c_xvs = [], [], [], []
        for layer in self.layers:
            x, a_xk, a_xv, c_xk, c_xv = layer(
                x, freqs_cos, freqs_sin, mask, encoder_out, encoder_mask
            )
            a_xks.append(a_xk)
            a_xvs.append(a_xv)
            c_xks.append(c_xk)
            c_xvs.append(c_xv)
        a_xks_stack = torch.stack(a_xks, dim=2)
        a_xvs_stack = torch.stack(a_xvs, dim=2)
        if self.layers[0].enable_cross_attention:
            return (
                x,
                a_xks_stack,
                a_xvs_stack,
                torch.stack(c_xks),
                torch.stack(c_xvs),
            )
        else:
            return x, a_xks_stack, a_xvs_stack


class TransformerBlockSequenceOneStep(nn.Module):
    def __init__(self, module_list):
        super().__init__()
        self.layers = module_list

    def forward(
        self,
        x: torch.Tensor,
        a_xks: torch.Tensor,
        a_xvs: torch.Tensor,
        c_xks: Optional[torch.Tensor],
        c_xvs: Optional[torch.Tensor],
        freqs_cos: torch.Tensor,
        freqs_sin: torch.Tensor,
        mask: torch.Tensor,
        encoder_mask: Optional[torch.Tensor],
    ):
        a_xk_news, a_xv_news = [], []
        for i, layer in enumerate(self.layers):
            (
                x,
                a_xk_new,
                a_xv_new,
            ) = layer(
                x,
                a_xks[:, :, i],
                a_xvs[:, :, i],
                c_xks[i] if c_xks is not None else None,
                c_xvs[i] if c_xvs is not None else None,
                freqs_cos,
                freqs_sin,
                mask,
                encoder_mask,
            )
            a_xk_news.append(a_xk_new)
            a_xv_news.append(a_xv_new)
        return (
            x,
            torch.stack(a_xk_news, dim=2),
            torch.stack(a_xv_news, dim=2),
        )


class Transformer(nn.Module):
    def __init__(self, params: ModelArgs):
        super().__init__()
        self.params = params
        self.vocab_size = params.vocab_size
        self.n_layers = params.n_layers

        self.tok_embeddings = nn.Embedding(params.vocab_size, params.dim)

        layers = torch.nn.ModuleList()
        for layer_id in range(params.n_layers):
            layers.append(TransformerBlock(layer_id, params))
        self.layer_sequence = TransformerBlockSequence(layers)

        self.norm = RMSNorm(params.dim, eps=params.norm_eps)
        self.output = nn.Linear(params.dim, params.vocab_size, bias=False)

        self.freqs_cos, self.freqs_sin = precompute_freqs_cos_sins(
            self.params.dim // self.params.n_heads,
            self.params.max_inference_seq_len * 2,
        )

        self.enable_cross_attention = params.enable_cross_attention

    def forward(
        self,
        tokens: torch.Tensor,
        aux_embeddings: torch.Tensor,
        encoder_out: Optional[
            torch.Tensor
        ] = None,  # expected (ebsz, encoder_seqlen, encoder_dim)
        encoder_valid: Optional[
            torch.Tensor
        ] = None,  # expected (ebsz, encoder_seqlen) of bools
    ):
        device = tokens.device
        h = self.tok_embeddings(tokens)
        self.freqs_cos = self.freqs_cos.to(device)
        self.freqs_sin = self.freqs_sin.to(device)

        h += aux_embeddings.to(device)

        seqlen = h.shape[1]

        mask = torch.full((seqlen, seqlen), float("-inf"), device=device)
        mask.triu_(diagonal=1)
        mask = mask.type_as(h)

        # encoder_mask should be (ebsz, 1, seqlen, encoder_seqlen)
        # where encoder_mask[..., i, j] = 0 if decoder token i can attend to encoder token j, else -inf
        if encoder_out is not None:
            assert encoder_valid is not None
            ebsz, encoder_seqlen = encoder_valid.shape
            assert (
                encoder_out.shape[0] == ebsz and encoder_out.shape[1] == encoder_seqlen
            )
            encoder_mask = torch.full(
                (ebsz, seqlen, encoder_seqlen), float("-inf"), device=device
            )
            encoder_mask.masked_fill_(encoder_valid.unsqueeze(1), 0)
            encoder_mask = encoder_mask.unsqueeze(1)
        else:
            encoder_mask = None

        h, *cached_vals = self.layer_sequence(
            h,
            self.freqs_cos,
            self.freqs_sin,
            mask,
            encoder_out,
            encoder_mask,
        )

        h = self.norm(h)

        # return the embeddings for the last non-pad token in each sequence
        # TODO update this
        idxs = (
            torch.argmax((tokens == self.params.padding_idx).to(torch.int32), dim=-1)
            - 1
        )
        idxs[idxs == -1] = tokens.shape[1] - 1
        embeddings = h[torch.arange(h.shape[0], dtype=torch.int32), idxs]

        return (
            F.log_softmax(self.output(h), dim=-1),
            *cached_vals,
            embeddings,
        )


class TransformerOneStep(Transformer):
    def __init__(self, params: ModelArgs):
        super().__init__(params)
        layers = torch.nn.ModuleList()
        for layer_id in range(params.n_layers):
            layers.append(TransformerBlockOneStep(layer_id, params))
        self.layer_sequence = TransformerBlockSequenceOneStep(layers)

    def forward(
        self,
        tokens: torch.Tensor,
        aux_embeddings: torch.Tensor,
        a_xks: torch.Tensor,
        a_xvs: torch.Tensor,
        c_xks: Optional[torch.Tensor],
        c_xvs: Optional[torch.Tensor],
        encoder_valid: Optional[
            torch.Tensor
        ],  # expected (ebsz, encoder_seqlen) of bools
    ):
        start_pos, _, _, _, _ = a_xks.shape
        seqlen = tokens.shape[1]
        h = self.tok_embeddings(tokens)
        self.freqs_cos = self.freqs_cos.to(h.device)
        self.freqs_sin = self.freqs_sin.to(h.device)

        h += aux_embeddings.to(h.device)

        mask = torch.full(
            (seqlen, start_pos + seqlen), float("-inf"), device=tokens.device
        )
        mask.triu_(diagonal=1 + start_pos)
        mask = mask.type_as(h)

        # encoder_mask should be (ebsz, 1, new_seqlen, encoder_seqlen)
        # where encoder_mask[..., i, j] = 0 if decoder token i can attend to encoder token j, else -inf
        if c_xvs is not None:
            assert encoder_valid is not None
            ebsz, encoder_seqlen = encoder_valid.shape
            assert c_xvs.shape[1] == ebsz and c_xvs.shape[3] == encoder_seqlen
            encoder_mask = torch.full(
                (ebsz, seqlen, encoder_seqlen), float("-inf"), device=tokens.device
            )
            encoder_mask.masked_fill_(encoder_valid.unsqueeze(1), 0)
            encoder_mask = encoder_mask.unsqueeze(1)
        else:
            encoder_mask = None

        h, a_xk_news, a_xv_news = self.layer_sequence(
            h,
            a_xks,
            a_xvs,
            c_xks,
            c_xvs,
            self.freqs_cos,
            self.freqs_sin,
            mask,
            encoder_mask,
        )

        h = self.norm(h)

        return F.log_softmax(self.output(h), dim=-1), a_xk_news, a_xv_news


class MusicalPositionEmbedTransformer(nn.Module):
    def __init__(self, vocab, params):
        super().__init__()

        # ONNX doesn't support flash attention
        assert not params.enable_flash

        self.vocab = vocab
        self.params = params
        self.transformer = Transformer(params)
        beats_max = vocab.embed_length_max // vocab.quantize_divisions
        self.bar_embedding = nn.Embedding(
            beats_max // 4 + 1, params.context_embedding_dim
        )
        self.beat_embedding = nn.Embedding(4, params.context_embedding_dim)
        self.tick_embedding = nn.Embedding(
            vocab.quantize_divisions, params.context_embedding_dim
        )
        self.polyphony_embedding = nn.Embedding(
            vocab.embed_polyphony_max + 1, params.context_embedding_dim
        )
        self.context_mix = nn.Linear(
            params.context_embedding_dim * 8, params.dim, bias=False
        )
        self.context_norm = RMSNorm(params.dim)

        self.param_count = 0
        for p in self.parameters():
            self.param_count += p.numel()

        self.dim = params.dim

    def get_context(self, x):
        context = self.context_mix(
            torch.cat(
                [
                    self.bar_embedding(x[..., 1]),
                    self.beat_embedding(x[..., 2]),
                    self.tick_embedding(x[..., 3]),
                    self.polyphony_embedding(x[..., 4]),
                    self.bar_embedding(x[..., 5]),
                    self.beat_embedding(x[..., 6]),
                    self.tick_embedding(x[..., 7]),
                    self.polyphony_embedding(x[..., 8]),
                ],
                dim=-1,
            )
        )
        return self.context_norm(context)

    def forward(
        self,
        x: torch.Tensor,
        encoder_out: Optional[torch.Tensor] = None,
        encoder_valid: Optional[torch.Tensor] = None,
    ):
        return self.transformer(
            x[..., 0],
            aux_embeddings=self.get_context(x),
            encoder_out=encoder_out,
            encoder_valid=encoder_valid,
        )


class MusicalPositionEmbedTransformerOneStep(MusicalPositionEmbedTransformer):
    def __init__(self, vocab, params):
        super().__init__(vocab, params)
        self.transformer = TransformerOneStep(params)

    def forward(
        self,
        x: torch.Tensor,
        a_xks: torch.Tensor,
        a_xvs: torch.Tensor,
        c_xks: Optional[torch.Tensor] = None,
        c_xvs: Optional[torch.Tensor] = None,
        encoder_valid: Optional[torch.Tensor] = None,
    ):
        return self.transformer(
            x[..., 0],
            aux_embeddings=self.get_context(x),
            a_xks=a_xks,
            a_xvs=a_xvs,
            c_xks=c_xks,
            c_xvs=c_xvs,
            encoder_valid=encoder_valid,
        )


class MusicalPositionEmbedTransformerOneStepNoCross(
    MusicalPositionEmbedTransformerOneStep
):
    def __init__(self, vocab, params):
        super().__init__(vocab, params)

    def forward(
        self,
        x: torch.Tensor,
        a_xks: torch.Tensor,
        a_xvs: torch.Tensor,
    ):
        return super().forward(x, a_xks, a_xvs)
