()
| 127 | model.config.use_cache = use_cache |
| 128 | |
| 129 | def main(): |
| 130 | import argparse |
| 131 | parser = argparse.ArgumentParser() |
| 132 | parser.add_argument('--model', default='meta-llama/Llama-2-7b-hf', type=str) # quant base model |
| 133 | parser.add_argument('--dev', type=str, default="cuda:0") |
| 134 | parser.add_argument('--quant_type', type=str, default="int", help='Quantization data type') |
| 135 | parser.add_argument('--bits', type=int, default=3, help='Quantization bits') |
| 136 | parser.add_argument('--group_size', type=int, default=128, help='Quantization group size') |
| 137 | |
| 138 | args = parser.parse_args() |
| 139 | print(args) |
| 140 | |
| 141 | |
| 142 | print("loading the model...") |
| 143 | model = AutoModelForCausalLM.from_pretrained(args.model, torch_dtype=torch.bfloat16, use_safetensors=True, low_cpu_mem_usage=True) |
| 144 | |
| 145 | q_config = { |
| 146 | "zero_point": True, # by default True |
| 147 | "q_group_size": args.group_size, # whether to use group quantization |
| 148 | } |
| 149 | model = model.cuda() |
| 150 | pseudo_quantize_model_weight( |
| 151 | model, w_bit=args.bits, q_config=q_config, quant_type=args.quant_type |
| 152 | ) |
| 153 | |
| 154 | dev = torch.device(args.dev) |
| 155 | |
| 156 | dataloader, testloader = get_wikitext2(nsamples=128, seed=0, seqlen=2048, model=args.model) |
| 157 | |
| 158 | llama_eval(model, testloader, dev) |
| 159 | |
| 160 | if __name__ == "__main__": |
| 161 | import logging |
no test coverage detected