Matir/VibeVoice-ASR-HF
029
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 