import argparse import torch from peft import PeftModel from transformers import AutoModelForCausalLM, AutoTokenizer from config import MODEL_ID, OUTPUT_DIR, SYSTEM_PROMPT def load_model(adapter_path=OUTPUT_DIR): use_cuda = torch.cuda.is_available() dtype = ( torch.bfloat16 if use_cuda and torch.cuda.is_bf16_supported() else torch.float16 if use_cuda else torch.float32 ) tokenizer = AutoTokenizer.from_pretrained(adapter_path) base_model = AutoModelForCausalLM.from_pretrained( MODEL_ID, torch_dtype=dtype, device_map="auto" if use_cuda else None, ) model = PeftModel.from_pretrained(base_model, adapter_path) model.eval() return model, tokenizer def generate_sql(model, tokenizer, sql_context, sql_prompt, max_new_tokens=256): user_message = ( "Database context:\n" f"{sql_context}\n\n" "Request:\n" f"{sql_prompt}" ) messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_message}, ] text = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, ) inputs = tokenizer(text, return_tensors="pt") device = next(model.parameters()).device inputs = {k: v.to(device) for k, v in inputs.items()} with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=False, pad_token_id=tokenizer.eos_token_id, ) generated_tokens = outputs[0][inputs["input_ids"].shape[1] :] return tokenizer.decode(generated_tokens, skip_special_tokens=True).strip() def main(): parser = argparse.ArgumentParser() parser.add_argument("--schema", required=True, help="Database schema/context") parser.add_argument("--question", required=True, help="Natural-language SQL request") parser.add_argument("--adapter", default=OUTPUT_DIR, help="Path to trained LoRA adapter") args = parser.parse_args() model, tokenizer = load_model(args.adapter) sql = generate_sql(model, tokenizer, args.schema, args.question) print("\nGenerated SQL:\n") print(sql) if __name__ == "__main__": main()