Harrison/pretty print (#57)

make stuff look nice
This commit is contained in:
Harrison Chase
2022-11-03 00:41:07 -07:00
committed by GitHub
parent dfb81c969f
commit 4cc18d6c2a
11 changed files with 180 additions and 64 deletions
+12 -6
View File
@@ -6,6 +6,7 @@ from pydantic import BaseModel, Extra
from langchain.chains.base import Chain
from langchain.chains.llm import LLMChain
from langchain.chains.sql_database.prompt import PROMPT
from langchain.input import ChainedInput
from langchain.llms.base import LLM
from langchain.sql_database import SQLDatabase
@@ -25,6 +26,7 @@ class SQLDatabaseChain(Chain, BaseModel):
"""LLM wrapper to use."""
database: SQLDatabase
"""SQL Database to connect to."""
verbose: bool = False
input_key: str = "query" #: :meta private:
output_key: str = "result" #: :meta private:
@@ -52,20 +54,24 @@ class SQLDatabaseChain(Chain, BaseModel):
def _run(self, inputs: Dict[str, str]) -> Dict[str, str]:
llm_chain = LLMChain(llm=self.llm, prompt=PROMPT)
_input = inputs[self.input_key] + "\nSQLQuery:"
chained_input = ChainedInput(
inputs[self.input_key] + "\nSQLQuery:", verbose=self.verbose
)
llm_inputs = {
"input": _input,
"input": chained_input.input,
"dialect": self.database.dialect,
"table_info": self.database.table_info,
"stop": ["\nSQLResult:"],
}
sql_cmd = llm_chain.predict(**llm_inputs)
print(sql_cmd)
chained_input.add(sql_cmd, color="green")
result = self.database.run(sql_cmd)
print(result)
_input += f"\nSQLResult: {result}\nAnswer:"
llm_inputs["input"] = _input
chained_input.add("\nSQLResult: ")
chained_input.add(result, color="yellow")
chained_input.add("\nAnswer:")
llm_inputs["input"] = chained_input.input
final_result = llm_chain.predict(**llm_inputs)
chained_input.add(final_result, color="green")
return {self.output_key: final_result}
def query(self, query: str) -> str: