AD1TEYA/multiscale-keypoint-conditioned-LipReading
8
1import logging2from argparse import ArgumentParser3 4import torch5import torchaudio6from datamodule.data_module import DataModule7from pytorch_lightning import Trainer8 9 10# Set environment variables and logger level11logging.basicConfig(level=logging.WARNING)12 13 14def get_trainer(args):15 return Trainer(num_nodes=1, devices=1, accelerator="gpu")16 17 18def get_lightning_module(args):19 # Set modules and trainer20 from lightning import ModelModule21 modelmodule = ModelModule(args)22 return modelmodule23 24 25def parse_args():26 parser = ArgumentParser()27 parser.add_argument(28 "--modality",29 type=str,30 help="Type of input modality",31 required=True,32 choices=["audio", "video"],33 )34 parser.add_argument(35 "--root-dir",36 type=str,37 help="Root directory of preprocessed dataset",38 required=True,39 )40 parser.add_argument(41 "--test-file",42 default="lrs3_test_transcript_lengths_seg16s.csv",43 type=str,44 help="Filename of testing label list. (Default: lrs3_test_transcript_lengths_seg16s.csv)",45 required=True,46 )47 parser.add_argument(48 "--pretrained-model-path",49 type=str,50 help="Path to the pre-trained model",51 required=True,52 )53 parser.add_argument(54 "--decode-snr-target",55 type=float,56 default=999999,57 help="Level of signal-to-noise ratio (SNR)",58 )59 parser.add_argument(60 "--debug",61 action="store_true",62 help="Flag to use debug level for logging",63 )64 return parser.parse_args()65 66 67def init_logger(debug):68 fmt = "%(asctime)s %(message)s" if debug else "%(message)s"69 level = logging.DEBUG if debug else logging.INFO70 logging.basicConfig(format=fmt, level=level, datefmt="%Y-%m-%d %H:%M:%S")71 72 73def cli_main():74 args = parse_args()75 init_logger(args.debug)76 modelmodule = get_lightning_module(args)77 datamodule = DataModule(args)78 trainer = get_trainer(args)79 trainer.test(model=modelmodule, datamodule=datamodule)80 81 82if __name__ == "__main__":83 cli_main()84 