import torch
import poptorch
import torchvision
import torch.nn as nn
import flash_attention_ipu.auto
import math
import torch.nn.functional as F

class MLPBlock(nn.Module):
    def __init__(self, emb_dim, mlp_dim):
        super().__init__()
        self.emb_dim = emb_dim
        self.mlp_dim = mlp_dim
        self.fc1 = nn.Linear(emb_dim, mlp_dim)
        self.fc2 = nn.Linear(mlp_dim, emb_dim)

    def forward(self, x, res):
        x = F.gelu(self.fc1(x))
        x = self.fc2(x) + res
        return x

class Encoder1DBlock(nn.Module):
    def __init__(self, emb_dim, mlp_dim, num_heads, batch_size, seq_len):
        super().__init__()
        self.ln1 = nn.LayerNorm(emb_dim)
        self.qkv = nn.Linear(emb_dim, emb_dim * 3, bias=False)
        self.attn_proj = nn.Linear(emb_dim, emb_dim)
        self.head_dim = emb_dim // num_heads
        self.mlp = MLPBlock(emb_dim, mlp_dim)
        self.ln2 = nn.LayerNorm(emb_dim)
        self.emb_dim = emb_dim
        self.num_heads = num_heads

    def forward(self, x):
        B, S, D = x.shape
        qkv = self.qkv(self.ln1(x))
        q, k, v = torch.split(qkv, self.emb_dim, dim=-1)
        attn_result = F.scaled_dot_product_attention(
            q.reshape(B, S, self.num_heads, self.head_dim), k.reshape(B, S, self.num_heads, self.head_dim), v.reshape(B, S, self.num_heads, self.head_dim), attn_mask=None, dropout_p=0.0, is_causal=True # TODO fix backend
        ).reshape(B, S, D)
        x = self.attn_proj(attn_result) + x
        x = self.mlp(self.ln2(x), x)
        return x

class Encoder(nn.Module):
    def __init__(self, emb_dim, mlp_dim, num_heads, batch_size, seq_len, depth):
        super().__init__()
        self.layers = nn.ModuleList([ Encoder1DBlock(emb_dim, mlp_dim, num_heads, batch_size, seq_len) for i in range(depth) ])
        self.ln = nn.LayerNorm(emb_dim)

    def forward(self, x):
        for layer in self.layers:
            x = layer(x)
        return self.ln(x)

class PositionalEmbeddings(nn.Module):
    def __init__(self, emb_dim, seq_len):
        super().__init__()
        self.pos_emb = nn.Parameter(torch.zeros(size=(1, seq_len, emb_dim), dtype=torch.float16))

    def forward(self, x):
        return x + self.pos_emb

class PatchEmbedder(nn.Module):
    def __init__(self, img_size, patch_size, in_chans, emb_dim):
        super().__init__()
        self.proj = nn.Conv2d(in_chans, emb_dim, kernel_size=patch_size, stride=patch_size, padding=0)
        self.flatten = True
        self.emb_dim = emb_dim
        self.proj_norm = nn.Identity()

    def forward(self, x):
        B, C, H, W = x.shape
        x = self.proj(x)
        if self.flatten:
            x = x.reshape(B, -1, self.emb_dim)
        x = self.proj_norm(x)
        return x

class MAPHead(nn.Module):
    def __init__(self, emb_dim, mlp_dim, num_heads, batch_size, seq_len):
        super().__init__()
        self.q = nn.Linear(emb_dim, emb_dim)
        self.kv = nn.Linear(emb_dim, emb_dim * 2)
        self.num_heads = num_heads
        self.head_dim = emb_dim // num_heads
        #self.q_norm = nn.LayerNorm(self.head_dim)
        #self.k_norm = nn.LayerNorm(self.head_dim)
        self.proj = nn.Linear(emb_dim, emb_dim)
        self.ln = nn.LayerNorm(emb_dim)
        self.mlp = MLPBlock(emb_dim, mlp_dim)
        self.probe = nn.Parameter(torch.zeros(size=[1, 1, emb_dim], dtype=torch.float16))
        self.batch_size = batch_size
        self.seq_len = seq_len
        self.emb_dim = emb_dim

    def forward(self, x):
        ql = self.probe.expand(self.batch_size, -1, -1)
        q = self.q(ql).reshape([self.batch_size, self.num_heads, 1, self.head_dim])
        kv = self.kv(x).reshape([self.batch_size, self.seq_len, 2, self.num_heads, self.head_dim]).permute((2, 0, 3, 1, 4))
        k, v = kv.split([1, 1], dim=0)
        k, v = k.squeeze(0), v.squeeze(0)
        #q = self.q_norm(q)
        #k = self.k_norm(k)
        #x = F.scaled_dot_product_attention(q, k, v, is_causal=True) # TODO fix causality
        # compute attention manually because bad (lash_attention_ipu does not currently support Grouped- or Multi-query attention)
        xr = F.softmax(torch.matmul(q, k.mT) / math.sqrt(self.head_dim), dim=-1)
        x = xr @ v
        x = x.transpose(1, 2).reshape((self.batch_size, 1, self.emb_dim))
        x = self.proj(x)
        return self.mlp(self.ln(x), x)

class VisionTransformer(nn.Module):
    def __init__(self, emb_dim, mlp_dim, num_heads, batch_size, seq_len, depth, img_size, patch_size, in_chans):
        super().__init__()
        self.patch_embed = PatchEmbedder(img_size, patch_size, in_chans, emb_dim)
        self.encoder = Encoder(emb_dim, mlp_dim, num_heads, batch_size, seq_len, depth)
        self.pool = MAPHead(emb_dim, mlp_dim, num_heads, batch_size, seq_len)
        self.pos_emb = PositionalEmbeddings(emb_dim, seq_len)

    def forward(self, image):
        x = self.pos_emb(self.patch_embed(image))
        return self.pool(self.encoder(x))

siglip_so400m_384_14 = {
    "img_size": 384, # 384/14 is ~27.4
    "emb_dim": 1152,
    "depth": 27,
    "num_heads": 16,
    "mlp_dim": 4304,
    "patch_size": 14,
    "in_chans": 3
}

batch_size = 4
seq_len = 729

model = VisionTransformer(**{
    **siglip_so400m_384_14,
    "batch_size": batch_size,
    "seq_len": seq_len
})
model.eval().half()
opts = poptorch.Options()
poptorch_model_inf = poptorch.inferenceModel(model, options=opts)
res = poptorch_model_inf(torch.randn(batch_size, 3, siglip_so400m_384_14["img_size"], siglip_so400m_384_14["img_size"]).half())
print(res)
