dvwn/nl2sql-api
0
1# Path: src/nl2sql/hf_engine.py2# This module defines the HuggingFace-based engine for generating SQL queries from natural language questions.3import os4from huggingface_hub import InferenceClient5from langchain_huggingface import HuggingFaceEndpoint6from langchain_core.language_models.llms import LLM7from typing import Any, List, Optional8 9# Model Registry: Add several model to be tested10MODEL_REGISTRY = {11 "defog/sqlcoder-7b-2": "text",12 "Qwen/Qwen2.5-Coder-7B-Instruct:featherless-ai": "chat",13 "Qwen/Qwen2.5-Coder-32B-Instruct:featherless-ai": "chat",14 "defog/llama-3-sqlcoder-8b:featherless-ai": "chat"15 #"deepseek-ai/DeepSeek-R1-Distill-Qwen-32B:featherless-ai": "chat"16}17 18DEFAULT_MODEL_ID = "Qwen/Qwen2.5-Coder-7B-Instruct:featherless-ai"19 20# Custom LangChain wrapper for HuggingFace Inference API21class HFChatWrapper(LLM):22 """23 Custom LLM wrapper for HuggingFace Inference API to maintain compatibility with LangChain's LLM interface.24 """25 client: Any26 model_id: str27 28 def _call(self, prompt: str, stop: Optional[List[str]] = None) -> str:29 completion = self.client.chat.completions.create(30 model = self.model_id,31 messages = [32 {"role": "user", "content": prompt}33 ],34 temperature = 0.0,35 max_tokens = 51236 )37 return completion.choices[0].message.content38 39 @property40 def _llm_type(self) -> str:41 return "huggingface_inference_client"42 43def get_models() -> List[str]:44 """Utility to return all model IDs available in the MODEL_REGISTRY."""45 return list(MODEL_REGISTRY.keys())46 47# Initialize the HuggingFace endpoint using the InferenceClient48def get_llm(model_id: str = DEFAULT_MODEL_ID):49 """50 Automatically detects the model type and returns the correct LangChain interface.51 Initializes the HuggingFace InferenceClient and returns an LLM instance for generating SQL queries.52 """53 # Load HuggingFace API token from environment variable54 hf_token = os.getenv("HF_TOKEN")55 if not hf_token:56 raise ValueError("HuggingFace API token not found!")57 58 # Determine the model type based on the MODEL_REGISTRY59 active_model = model_id if model_id else DEFAULT_MODEL_ID60 61 if active_model not in MODEL_REGISTRY:62 print(f"Warning: Model '{active_model}' not found in MODEL_REGISTRY. Defaulting to 'chat' type.")63 64 model_type = MODEL_REGISTRY.get(active_model, "chat")65 print(f"Initializing HuggingFace InferenceClient with model: {active_model}")66 67 if model_type == "chat":68 client = InferenceClient(api_key=hf_token)69 return HFChatWrapper(client=client, model_id=active_model)70 elif model_type == "text":71 # Route to standard Text Generation API72 return HuggingFaceEndpoint(73 repo_id=active_model,74 task="text-generation",75 max_new_tokens=512,76 temperature=0.0,77 huggingfacehub_api_token=hf_token,78 do_sample=False,79 return_full_text=False80 )81 else:82 raise ValueError(f"Unknown model type: {model_type}")83 84 # Initialize the HuggingFace InferenceClient85 #client = InferenceClient(api_key=hf_token)86 #llm = HFChatWrapper(client=client, model_id=active_model)87 88 #return llm89 90if __name__=="__main__":91 from dotenv import load_dotenv92 load_dotenv()93 94 try:95 test_llm = get_llm()96 print("Model loaded successfully! Running a quick ping...")97 response = test_llm.invoke("write a single SQL statement to count all rows in a table name 'Employee'.")98 print(f"\nResponse:\n{response}")99 except Exception as e:100 print(f"Error during LLM initialization: {e}")