CoolFace
Modelpublic

Matir/VibeVoice-ASR-HF

sourceHugging Facemitupdated 3mo agoView on Hugging Face
0likes29downloads
quantize.py114 linesDownload Raw Back to root
1#!/usr/bin/env python32import argparse3import os4import sys5import logging6import torch7from transformers import AutoProcessor, BitsAndBytesConfig, TorchAoConfig8 9# We need to bypass lazy-loader for VibeVoice classes as in main.py10try:11    from transformers.models.vibevoice_asr.modeling_vibevoice_asr import VibeVoiceAsrForConditionalGeneration12except ImportError as e:13    print(f"Error importing VibeVoice modeling: {e}", file=sys.stderr)14    print("Please ensure the correct transformers version is installed.", file=sys.stderr)15    sys.exit(1)16 17logging.basicConfig(18    level=logging.INFO,19    format="%(asctime)s [%(levelname)s] %(message)s",20    datefmt="%Y-%m-%d %H:%M:%S"21)22logger = logging.getLogger(__name__)23 24def main():25    parser = argparse.ArgumentParser(description="Quantize VibeVoice ASR model and save the serialized weights.")26    parser.add_argument("--model_dir", type=str, default=os.environ.get("MODEL_DIR", "./repository"),27                        help="Path to the original unquantized model directory (default: ./repository or $MODEL_DIR)")28    parser.add_argument("--output_dir", type=str, required=True,29                        help="Path where the quantized model should be saved (must be a new/empty directory)")30    parser.add_argument("--format", type=str, choices=["int8", "int4", "fp8"], default="int8",31                        help="Quantization format (int8, int4, or fp8, default: int8)")32 33    args = parser.parse_args()34 35    if not os.path.exists(args.model_dir):36        logger.error(f"Source model directory does not exist: {args.model_dir}")37        sys.exit(1)38 39    # Clean path resolution40    args.model_dir = os.path.abspath(args.model_dir)41    args.output_dir = os.path.abspath(args.output_dir)42 43    if os.path.exists(args.output_dir) and os.path.exists(os.path.join(args.output_dir, "config.json")):44        logger.warning(f"Output directory '{args.output_dir}' already contains a model config. "45                       f"Saving here will overwrite it and may leave orphan weight files.")46        confirm = input("Do you want to proceed anyway? (y/N): ")47        if confirm.lower() != 'y':48            logger.info("Aborted by user.")49            sys.exit(0)50 51    # Configure quantization52    if args.format == "int8":53        logger.info("Configuring 8-bit Integer (BitsAndBytes) quantization...")54        quantization_config = BitsAndBytesConfig(load_in_8bit=True)55    elif args.format == "int4":56        logger.info("Configuring 4-bit (NF4 BitsAndBytes) quantization...")57        quantization_config = BitsAndBytesConfig(58            load_in_4bit=True,59            bnb_4bit_compute_dtype=torch.bfloat16,60            bnb_4bit_use_double_quant=True,61            bnb_4bit_quant_type="nf4"62        )63    elif args.format == "fp8":64        logger.info("Configuring 8-bit Floating Point (TorchAO Float8 Weight-Only) quantization...")65        try:66            from torchao.quantization import Float8WeightOnlyConfig67        except ImportError:68            logger.error("torchao is not installed. Please run 'pip install torchao' to use FP8 quantization.")69            sys.exit(1)70        quantization_config = TorchAoConfig(Float8WeightOnlyConfig())71 72    # Check CUDA availability73    if not torch.cuda.is_available():74        logger.error("CUDA is not available. BitsAndBytes quantization requires a GPU to perform the compression.")75        sys.exit(1)76 77    logger.info(f"Loading model from '{args.model_dir}' and quantizing to {args.format}...")78    try:79        # Load config and force Flash Attention 2 on the text decoder80        from transformers import AutoConfig81        config = AutoConfig.from_pretrained(args.model_dir)82        if hasattr(config, "text_config"):83            config.text_config._attn_implementation = "flash_attention_2"84            logger.info("Forced Flash Attention 2 on the text decoder.")85 86        # Load model87        model = VibeVoiceAsrForConditionalGeneration.from_pretrained(88            args.model_dir,89            config=config,90            quantization_config=quantization_config,91            torch_dtype=torch.bfloat16,92            device_map="auto",93        )94 95        # Load processor96        logger.info("Loading processor...")97        processor = AutoProcessor.from_pretrained(args.model_dir)98 99        # Save100        logger.info(f"Saving quantized model and processor to '{args.output_dir}'...")101        os.makedirs(args.output_dir, exist_ok=True)102        model.save_pretrained(args.output_dir)103        processor.save_pretrained(args.output_dir)104 105        logger.info("Quantization completed successfully!")106        logger.info(f"You can now point your FastAPI server to '{args.output_dir}' to load it instantly.")107 108    except Exception as e:109        logger.exception(f"Quantization failed: {e}")110        sys.exit(1)111 112if __name__ == "__main__":113    main()114