# NeMo's "core" package
import nemo
import datetime

# NeMo's ASR collection - this collections contains complete ASR models and
# building blocks (modules) for ASR
import nemo.collections.asr as nemo_asr
from omegaconf import OmegaConf, open_dict
from pytorch_lightning.loggers import WandbLogger
from pytorch_lightning.callbacks import ModelCheckpoint
import pytorch_lightning as pl
import torch

torch.set_float32_matmul_precision("high")

# wandb_logger = None
wandb_logger = WandbLogger(
    log_model="all",
    project="hoot",
    name="t#est_en_fast",
    save_dir="/home/tony/Data/checkpoints",
)

params = OmegaConf.load("./configs/config_faster_conformer_bpe_test.yaml")

params.model.tokenizer.dir = "./tokenizers/en/tokenizer_spe_bpe_v1024/"  # note this is a directory, not a path to a vocabulary file
params.model.tokenizer.type = "bpe"

date_time_str = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
# trainer = pl.Trainer(devices=1, accelerator="gpu", max_epochs=50, logger=wandb_logger)
checkpoint_callback = ModelCheckpoint(
    dirpath=f"/home/tony/Data/checkpoints/{date_time_str}", every_n_train_steps=500
)
trainer = pl.Trainer(**params.trainer, callbacks=[checkpoint_callback])
trainer.logger = wandb_logger

train_manifest = "/home/tony/Data/Hoot/en_train_manifest_fast_longer.json"
test_manifest = "/home/tony/Data/Hoot/en_test_manifest_fast.json"

# Update paths to dataset
params.model.train_ds.manifest_filepath = train_manifest
params.model.validation_ds.manifest_filepath = test_manifest

first_asr_model = nemo_asr.models.EncDecCTCModelBPE(cfg=params.model, trainer=trainer)
preload_checkpoint = "/home/tony/Data/Hoot/stt_en_fastconformer_ctc_large.pt"
print("preloading checkpoint")
cur_state_dict = first_asr_model.state_dict()
checkpoint = torch.load(preload_checkpoint)
state_dict = checkpoint["model"] if "model" in checkpoint else checkpoint
print(
    "before loading", first_asr_model.encoder.layers[-1].self_attn.linear_q.weight.std()
)
first_asr_model.load_state_dict(state_dict, strict=False)
print(
    "after loading", first_asr_model.encoder.layers[-1].self_attn.linear_q.weight.std()
)

# Start training!!!
trainer.fit(first_asr_model)
date_time_str = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
first_asr_model.save_to(f"/home/tony/Data/checkpoints/{date_time_str}/model_test.nemo")
print("done!")
