MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / main

Function main

test/general/wiki_ppl.py:129–158  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

127 model.config.use_cache = use_cache
128
129def 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
160if __name__ == "__main__":
161 import logging

Callers 1

wiki_ppl.pyFile · 0.70

Calls 4

get_wikitext2Function · 0.85
llama_evalFunction · 0.85
deviceMethod · 0.45

Tested by

no test coverage detected