# -- Code Cell --
import pandas as pd

# -- Code Cell --
train=pd.read_csv("train_data.csv")
test=pd.read_csv("test_data.csv")
unlabeled = pd.read_csv("unlabeled_data.csv")

# -- Code Cell --
unlabeled["text"]

# -- Code Cell --
unlabeled["text"][0]

# -- Code Cell --
unlabeled_tokenized = []
for i in range(len(unlabeled)):
    unlabeled_tokenized.append(unlabeled["text"][i].split())

# -- Code Cell --
unlabeled_tokenized

# -- Code Cell --
from gensim.models import Word2Vec
w2v = Word2Vec(sentences=unlabeled_tokenized,vector_size=200,window=5,min_count=1,epochs=10)

# -- Code Cell --
train_text = train["text"]
train_text

# -- Code Cell --
embeddings={}
for idx,line in enumerate(train_text):
    embeddings[idx] = [w2v.wv[word]for word in line.split()]

# -- Code Cell --
from torch.utils.data import DataLoader,Dataset,random_split
import torch.nn as nn
import torch

# -- Code Cell --
from torch.nn.utils.rnn import pad_sequence
def collate_fn(batch):
    X, y = zip(*batch)
    X = pad_sequence(X, batch_first=True, padding_value=0.0)
    y = torch.tensor(y, dtype=torch.long)
    return X, y

# -- Code Cell --
class TrainDataset(Dataset):
    def __init__(self,path):
        self.df = pd.read_csv(path)
    def __len__(self):
        return len(self.df)
    def __getitem__(self, index):
        row = self.df.iloc[index]
        text = embeddings[index]
        label = row['label']
        text = torch.tensor(text)
        return text,label

# -- Code Cell --
train_ds = TrainDataset("./train_data.csv")
train_subset, val_subset = random_split(train_ds, [0.8,0.2])

# -- Code Cell --
train_loader = DataLoader(train_subset, batch_size=32, shuffle=True, num_workers=0,collate_fn=collate_fn)
val_loader = DataLoader(val_subset, batch_size=32, shuffle=True, num_workers=0,collate_fn=collate_fn)

# -- Code Cell --
len(embeddings[80])

# -- Code Cell --
class Model(nn.Module):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.lstm = nn.LSTM(200, 128, num_layers=2, bidirectional=True, dropout=0.3, batch_first=True)
        self.clsf = nn.Sequential(
            nn.Linear(128*2, 64),
            nn.ReLU(),
            nn.Linear(64,4)
        )
    def forward(self, x):
        lstm, _ = self.lstm(x)
        out = self.clsf(lstm[:, -1, :])
        return out

# -- Code Cell --
device = 'cuda'
model = Model()
model = model.to(device)
import torch.optim as optim
optimizer = optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.2)

# -- Code Cell --
from sklearn.metrics import f1_score

# -- Code Cell --
for epoch in range(80):
    model.train()
    for text,label in train_loader:
        text =text.to(device)
        label = label.to(device)
        optimizer.zero_grad()
        output = model(text)
        loss= criterion(output,label)
        loss.backward()
        optimizer.step()
    model.eval()
    all_preds = []
    all_labels = []
    for text,label in val_loader:
        text =text.to(device)
        label = label.to(device)
        output = model(text)
        pred = output.argmax(dim=1)
        all_preds.append(pred.cpu())
        all_labels.append(label.cpu())
    all_preds = torch.cat(all_preds).numpy()
    all_labels = torch.cat(all_labels).numpy()
    f1 = f1_score(all_labels,all_preds,average='macro')
    scheduler.step()
    print(f"epoch {epoch}, f1 {f1}")

# -- Code Cell --
test_text= test['text']

# -- Code Cell --
embeddings_test={}
for idx,line in enumerate(test_text):
    embeddings_test[idx] = [w2v.wv[word]for word in line.split()]

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self,path):
        self.df = pd.read_csv(path)
    def __len__(self):
        return len(self.df)
    def __getitem__(self, index):
        text = embeddings_test[index]
        text = torch.tensor(text)
        return text

# -- Code Cell --
def collate_fn_test(batch):
    return pad_sequence(batch, padding_value=0.0,batch_first=True)

# -- Code Cell --
test_ds = TestDataset("./test_data.csv")
test_loader = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=0, collate_fn=collate_fn_test)

# -- Code Cell --
preds = []
with torch.no_grad():
    for text in test_loader:
            text =text.to(device)
            output = model(text)
            pred = output.argmax(dim=1)
            preds.append(pred.cpu())
preds = torch.cat(preds).numpy()
preds

# -- Code Cell --
sub1 = pd.DataFrame({
    "subtaskID": [1],
    "datapointID": [0],
    "answer": [3529]
})

sub2 = pd.DataFrame({
    "subtaskID": 2,
    "datapointID": test['id'],
    "answer": preds
})

final = pd.concat([sub1, sub2], ignore_index=True)
final.to_csv("subs.csv",index=False)

# -- Code Cell --
