# -- Code Cell --
import pandas as pd
train = pd.read_csv('train_data.csv')
test = pd.read_csv('test_data.csv')

# -- Code Cell --
train.head()

# -- Code Cell --
train_texts= train['comment_text']

# -- Code Cell --
from nltk.tokenize import word_tokenize

# -- Code Cell --
train_toknes = []
for i in range(len(train_texts)):
    train_toknes.append(word_tokenize(train_texts[i].lower()))

# -- Code Cell --
test_texts= test['comment_text']
test_tokens = []
for i in range(len(test_texts)):
    test_tokens.append(word_tokenize(test_texts[i].lower()))

# -- Code Cell --
PADDING = 300

# -- Code Cell --
for i in range(len(train_toknes)):
    if len(train_toknes[i]) > PADDING:
        train_toknes[i] = train_toknes[i][:PADDING]
    else:
        for j in range(PADDING - len(train_toknes[i])):
            train_toknes[i].append("<PAD>")

# -- Code Cell --
for i in range(len(train_toknes)):
    if len(train_toknes[i]) > PADDING:
        train_toknes[i] = train_toknes[i][:PADDING]
    else:
        for j in range(PADDING - len(train_toknes[i])):
            train_toknes[i].append("<PAD>")

# -- Code Cell --
len(test_tokens[10])

# -- Code Cell --
from collections import Counter
all_words = [w for comment in train_toknes for w in comment if w != "<PAD>"]
freq = Counter(all_words)
VOCAB_SIZE = 40000
most_commong = freq.most_common(VOCAB_SIZE)
word2index = {"<PAD>":0, "<UNK>":1}
for i,(word,_) in enumerate(most_commong):
    word2index[word] = i+1
def encode(token_list):
    return [word2index.get(w,1)for w in token_list]
X_train = [encode(t) for t in train_toknes]

# -- Code Cell --
all_words_test= [w for comment in test_tokens for w in comment if w != "<PAD>"]
X_test = [encode(t) for t in test_tokens]

# -- Code Cell --
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
from torchvision import models, transforms
from PIL import Image
from torchvision.transforms import functional as TF

# -- Code Cell --
train.head()

# -- Code Cell --
all_lables = train[["toxic", "severe_toxic","obscene","insult"]].values
all_lables

# -- 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):
        text = torch.tensor(X_train[index], dtype=torch.long)
        label = torch.tensor(all_lables[index], dtype=torch.float32)
        return text,label

# -- Code Cell --
train_ds = TrainDataset("./train_data.csv")

# -- Code Cell --
from torch.utils.data import random_split
train_size = int(0.8*len(train_ds))
val_size = len(train_ds) - train_size
train_subset, val_subset = random_split(train_ds,[train_size,val_size])

# -- 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)
    return X, torch.stack(y)

# -- Code Cell --
train_dataloader = DataLoader(train_subset, shuffle=True, num_workers=0, batch_size=32, collate_fn=collate_fn)
valid_dataloader = DataLoader(val_subset, shuffle=False, num_workers=0, batch_size=32, collate_fn=collate_fn)

# -- Code Cell --
device = torch.device("cuda")

# -- Code Cell --
VOCAB_SIZE = len(word2index)
VOCAB_SIZE

# -- Code Cell --
class BILSTM(nn.Module):
    def __init__(self):
        super().__init__()
        self.emb = nn.Embedding(VOCAB_SIZE, 128)
        self.lstm = nn.LSTM(128,128, bidirectional=True, batch_first=True, num_layers=2, dropout=0.3)
        self.fc = nn.Linear(256,4)
    def forward(self,x):
        emb = self.emb(x)
        lstm, _ = self.lstm(emb)
        lstm = lstm.max(dim=1)[0]
        out = self.fc(lstm)
        return out

# -- Code Cell --
model = BILSTM()
model = model.to(device)
optimizer = optim.Adam(model.parameters(), lr = 1e-3)
criterion = nn.BCEWithLogitsLoss()
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)

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

# -- Code Cell --
for epoch in range(10):
    model.train()
    total_loss = 0
    for img, label in train_dataloader:
        img = img.to(device)
        label = label.to(device)
        optimizer.zero_grad()
        output=model(img)
        loss = criterion(output, label)
        loss.backward()
        optimizer.step()
        total_loss+=loss.item()
    model.eval()
    val_loss = 0
    all_labels = []
    all_preds = []
    with torch.no_grad():
        for img,label in valid_dataloader:
            img = img.to(device)
            label = label.to(device)
            output=model(img)
            loss = criterion(output, label)
            val_loss += loss
            all_labels.append(label.cpu())
            pred = (torch.sigmoid(output) > 0.5).int()
            all_preds.append(pred.cpu())
    all_labels = torch.cat(all_labels)
    all_preds = torch.cat(all_preds)
    f1  = f1_score(all_labels, all_preds, average="macro")
    print(f"EPOCH {epoch+1}, LOSS_TRAIN {total_loss/len(train_dataloader)}, LOSS_VAL {val_loss/len(valid_dataloader)}, F1 {f1}")

# -- 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 = torch.tensor(X_test[index], dtype=torch.long)
        return text

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

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

# -- Code Cell --
preds=[]
for text in test:
    with torch.no_grad():
        text = text.to(device)
        pred = model(text)
        pred = torch.sigmoid(pred)
        preds.append(pred.cpu())
        pred_final = torch.cat(preds)
pred_final

# -- Code Cell --
