allinaigc/text2sql
0
1"""21. 这里用公网Qwen API来代替本地的ChatGLM模型。为了在Huggingface上演示。31. 使用确定了列名作为SQL语句的变量名,可以有效解决模型生成的SQL语句中变量名准确的问题。4 5"""6##TODO: 7 8import requests9import os10from rich import print11import os12import sys13import time14import pandas as pd15import numpy as np16import sys17import time18from typing import Any19import requests20import csv21import os22from rich import print23import pandas24import io25from io import StringIO26import re27from langchain.llms.utils import enforce_stop_tokens28import json29from transformers import AutoModel, AutoTokenizer30import mdtex2html31import qwen_response32 33 34''' Start: Environment settings. '''35os.environ['SENTENCE_TRANSFORMERS_HOME'] = '/Users/yunshi/Downloads/chatGLM/My_LocalKB_Project/'36os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"37import torch38mps_device = torch.device("mps") ## 在mac机器上需要加上这句。必须要有这句,否则会报错。39 40### 在langchain中定义chatGLM作为LLM。41from typing import Any, List, Mapping, Optional42from langchain.callbacks.manager import CallbackManagerForLLMRun43from langchain.llms.base import LLM44from transformers import AutoTokenizer, AutoModel45# llm_filepath = str("/Users/yunshi/Downloads/chatGLM/ChatGLM3-6B/6B") ## 第三代chatGLM 6B W/ code-interpreter46 47 48# ## API模式启动ChatGLM49# ## 配置ChatGLM的类与后端api server对应。50# class ChatGLM(LLM):51# max_token: int = 204852# temperature: float = 0.153# top_p = 0.954# history = []55 56# def __init__(self):57# super().__init__()58 59# @property60# def _llm_type(self) -> str:61# return "ChatGLM"62 63# def _call(self, prompt: str, stop: Optional[List[str]] = None) -> str:64# # headers中添加上content-type这个参数,指定为json格式65# headers = {'Content-Type': 'application/json'}66# data=json.dumps({67# 'prompt':prompt,68# 'temperature':self.temperature,69# 'history':self.history,70# 'max_length':self.max_token71# })72# print("ChatGLM prompt:",prompt)73# # 调用api74# # response = requests.post("http://0.0.0.0:8000",headers=headers,data=data) ##working。75# response = requests.post("http://127.0.0.1:8000",headers=headers,data=data) ##working。76# print("ChatGLM resp:", response)77 78# if response.status_code!=200:79# return "查询结果错误"80# resp = response.json()81# if stop is not None:82# response = enforce_stop_tokens(response, stop)83# self.history = self.history+[[None, resp['response']]] ##original84# return resp['response'] ##original.85 86# llm = ChatGLM() ## 启动一个实例。orignal working。87# import asyncio88# llm = ChatGLM() ## 启动一个实例。89 90 91''' End: Environment settings. '''92 93### 我会用中文或者英文双引号(即:“ ”," ")来告知你变量的名称。 长度","宽度","价格","产品ID","比率","类别","*"94 95### 用ChatGLM构建一个只返回SQL语句的模型。96def main(prompt):97 full_reponse = []98 sys_prompt = """99 1. 你是一个将文字转换成SQL语句的人工智能。100 2. 你需要注意:你只需要用纯文本回复代码的内容,即你不允许回复代码以外的任何信息。101 3. SQL变量默认是中文,而且只能从如下的名称列表中选择,你不可以使用这些名字以外的变量名:"长度","宽度","价格","产品ID","比率","类别","*"102 4. 你不能写IF, THEN的SQL语句,需要使用CASE。103 5. 我需要你转换的文字如下:""" 104 105 total_prompt = sys_prompt + "在数据表格table01中," + prompt106 107 print('total prompt now:',total_prompt)108 # for response, history in chatglm.model.stream_chat(chatglm.tokenizer, query=str(prompt)): ## 这里保留了所有的chat history在input_prompt中。109 # for response, history in chatglm.model.stream_chat(chatglm.tokenizer, query=str(prompt)): ## 这里保留了所有的chat history在input_prompt中。110 # for response, history in chatglm.model.stream_chat(chatglm.tokenizer, query=str(total_prompt)): ## 这里保留了所有的chat history在input_prompt中。111 112 # # for response, history in chatglm.model.stream_chat(chatglm.tokenizer, query=str(input_prompt[-1][0])): ## 从用langchain的自定义方式来做。113 # # for response, history in chatglm.model.stream_chat(chatglm.tokenizer, query=str(input_prompt[-1][0]), history=input_prompt, max_length=max_tokens, top_p=top_p, temperature=temperature): ## 从用langchain的自定义方式来做。114 # if response != "<br>":115 # # print('response of model:', response)116 # # input_prompt[-1][1] = response ## working.117 # # input_prompt[-1][1] = response118 # # yield input_prompt119 # full_reponse.append(response)120 121 # ## 得到一个非stream格式的答复。非API模式。 122 # response, history = chatglm.model.chat(chatglm.tokenizer, query=str(total_prompt), temperature=0.1) ## 这里保留了所有的chat history在input_prompt中。123 124 ###TODO:API模式,需要先启动API服务器。125 # llm = ChatGLM() ##!! 重要说明:每次都需要实例化一次!!!否则会报错content error。实际上是应该在每次函数调用的时候都要实例化一次!126 # response = llm(total_prompt) ## 这里是本地的ChatGLM来作为大模型输出基座。127 128 129 response = qwen_response.call_with_messages(total_prompt)130 131 print('response of model:', response)132 133 ## 用regex来提取纯SQL语句。需要构建多个正则式pattern134 pattern_1 = r"(?:`sql\n|\n`)"135 pattern_2 = r"(?:```|``)"136 pattern_3 = r"(?s)(.*?SQL语句示例.*?:).*?\n"137 pattern_4 = r"(?:`{3}|`{2}|`)"138 # pattern_5 = r"[\u4e00-\u9FFF]" ## 匹配中文。139 # pattern_6 = r"^[\u4e00-\u9fa5]{5,}" ## 首行中包含5个中文汉字的。140 pattern_7 = r"^.{0,2}([\u4e00-\u9fa5]{5,}).*" ## 首行中包含5个中文汉字的。 141 pattern_8 = r'^"|"$' ## 去除一句话开始或者末尾的英文双引号 142 pattern_list = [pattern_1, pattern_2, pattern_3, pattern_4, pattern_7, pattern_8]143 144 ## 遍历所有的pattern,逐个去除。145 full_reponse = response146 for p in pattern_list:147 full_reponse = re.sub(p, "", full_reponse)148 # final_response = re.sub(pattern_1, "", response) ## 逐步匹配。149 # final_response = re.sub(pattern_1, "", response) ## 逐步匹配。150 # final_response = re.sub(pattern_2, "", final_response) ## 逐步匹配。151 152 return full_reponse153 154# prompt = "你给我一段复杂的SQL语句示例。"155# prompt = "你给我一段SQL语句,用来完成如下工作:查询年龄大于30岁,男性,收入超过2万元的员工。"156# res = main(prompt=prompt)157# print(res)158 