mlnsio/text2sql
0
1# Prediction interface for Cog ⚙️2# https://github.com/replicate/cog/blob/main/docs/python.md3 4from cog import BasePredictor, Input5import torch6from transformers import AutoTokenizer, AutoModelForCausalLM, pipeline7import argparse8 9 10class Predictor(BasePredictor):11 def setup(self) -> None:12 """Load the model into memory to make running multiple predictions efficient"""13 # self.model = torch.load("./weights.pth")14 model_name = "defog/sqlcoder-34b-alpha"15 self.tokenizer = AutoTokenizer.from_pretrained(model_name)16 self.model = AutoModelForCausalLM.from_pretrained(17 model_name,18 torch_dtype=torch.float16,19 device_map="auto",20 use_cache=True,21 offload_folder="./.cache",22 )23 24 def predict(25 self,26 prompt: str = Input(description="Prompt to generate from"),27 ) -> str:28 """Run a single prediction on the model"""29 # processed_input = preprocess(image)30 # output = self.model(processed_image, scale)31 # return postprocess(output)32 33 # make sure the model stops generating at triple ticks34 # eos_token_id = tokenizer.convert_tokens_to_ids(["```"])[0]35 eos_token_id = self.tokenizer.eos_token_id36 pipe = pipeline(37 "text-generation",38 model=self.model,39 tokenizer=self.tokenizer,40 max_length=300,41 do_sample=False,42 num_beams=5, # do beam search with 5 beams for high quality results43 )44 generated_query = (45 pipe(46 prompt,47 num_return_sequences=1,48 eos_token_id=eos_token_id,49 pad_token_id=eos_token_id,50 )[0]["generated_text"]51 .split("```sql")[-1]52 .split("```")[0]53 .split(";")[0]54 .strip()55 + ";"56 )57 return generated_query58 