# 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!")