Team Ai
Apppublic

dvwn/nl2sql-api

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
hf_engine.py100 linesDownload Raw Back to nl2sql
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}")