Spaces:
Running on Zero
Running on Zero
File size: 2,280 Bytes
e99ee9c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 | 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()
|