CoolFace
Apppublic

AD1TEYA/multiscale-keypoint-conditioned-LipReading

sourceHugging Facemitupdated 6mo agoView on Hugging Face
8likes
eval.py84 linesDownload Raw Back to root
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