maykcaldas/MAPI_LLM
3
1from mp_api.client import MPRester2from emmet.core.summary import HasProps3import openai4import langchain5from langchain import OpenAI6from langchain import agents7from langchain.agents import initialize_agent8from langchain.agents import Tool, tool9from langchain import LLMMathChain, SerpAPIWrapper10from gpt_index import GPTListIndex, GPTIndexMemory11from langchain import SerpAPIWrapper12from langchain.prompts.few_shot import FewShotPromptTemplate13from langchain.prompts.prompt import PromptTemplate14from langchain.vectorstores import FAISS, Chroma15from langchain.embeddings import OpenAIEmbeddings16from langchain.prompts.example_selector import (MaxMarginalRelevanceExampleSelector, 17 SemanticSimilarityExampleSelector)18import requests19from rdkit import Chem20import pandas as pd21import os22 23class MAPITools:24 def __init__(self):25 self.model = 'text-ada-001' #maybe change to gpt-4 when ready26 self.k=1027 28 def get_material_atoms(self, formula):29 '''Receives a material formula and returns the atoms symbols present in it separated by comma.'''30 import re31 pattern = re.compile(r"([A-Z][a-z]*)(\d*)")32 matches = pattern.findall(formula)33 atoms = []34 for m in matches:35 atom, count = m36 count = int(count) if count else 137 atoms.append((atom, count))38 return ",".join([a[0] for a in atoms])39 40 def check_prop_by_formula(self, formula):41 raise NotImplementedError('Should be implemented in children classes')42 43 def search_similars_by_atom(self, atoms):44 '''This function receives a string with the atoms separated by comma as input and returns a list of similar materials'''45 atoms = atoms.replace(" ", "")46 with MPRester(os.getenv("MAPI_API_KEY")) as mpr:47 docs = mpr.summary.search(elements=atoms.split(','), fields=["formula_pretty", self.prop])48 return docs49 50 def create_context_prompt(self, formula):51 raise NotImplementedError('Should be implemented in children classes')52 53 def LLM_predict(self, prompt):54 ''' This function receives a prompt generate with context by the create_context_prompt tool and request a completion to a language model. Then returns the completion'''55 llm = OpenAI(56 model_name=self.model,57 temperature=0.7,58 n=1,59 best_of=5,60 top_p=1.0,61 stop=["\n\n", "###", "#", "##"],62 # model_kwargs=kwargs,63 )64 return llm.generate([prompt]).generations[0][0].text65 66 def get_tools(self):67 return [68 Tool(69 name = "Get atoms in material",70 func = self.get_material_atoms,71 description = (72 "Receives a material formula and returns the atoms symbols present in it separated by comma."73 )74 ),75 Tool(76 name = f"Checks if material is {self.prop_name} by formula",77 func = self.check_prop_by_formula,78 description = (79 f"This functions searches in the material project's API for the formula and returns if it is {self.prop_name} or not."80 )81 ),82 # Tool(83 # name = "Search similar materials by atom",84 # func = self.search_similars_by_atom,85 # description = (86 # "This function receives a string with the atoms separated by comma as input and returns a list of similar materials."87 # )88 # ),89 Tool(90 name = f"Create {self.prop_name} context to LLM search",91 func = self.create_context_prompt,92 description = (93 f"This function received a material formula as input and create a prompt to be inputed in the LLM_predict tool to predict if the material is {self.prop_name}." 94 if isinstance(self, MAPI_class_tools) else95 f"This function received a material formula as input and create a prompt to be inputed in the LLM_predict tool to predict the {self.prop_name} of a material." 96 )97 ),98 Tool(name = "LLM predictiom",99 func = self.LLM_predict,100 description = (101 "This function receives a prompt generate with context by the create_context_prompt tool and request a completion to a language model. Then returns the completion"102 )103 )104 ]105 106class MAPI_class_tools(MAPITools):107 def __init__(self, prop, prop_name, p_label, n_label):108 super().__init__()109 self.prop = prop110 self.prop_name = prop_name111 self.p_label = p_label112 self.n_label = n_label113 114 def check_prop_by_formula(self, formula):115 f''' This functions searches in the material project's API for the formula and returns if it is {self.prop_name} or not'''116 with MPRester(os.getenv("MAPI_API_KEY")) as mpr:117 docs = mpr.summary.search(formula=formula, fields=["formula_pretty", self.prop])118 if docs:119 if docs[0].formula_pretty == formula:120 return self.p_label if docs[0].dict()[self.prop] else self.n_label121 return f"Could not find any material while searching {formula}"122 123 def create_context_prompt(self, formula):124 '''This function received a material formula as input and create a prompt to be inputed in the LLM_predict tool to predict if the formula is a stable material '''125 elements = self.get_material_atoms(formula)126 similars = self.search_similars_by_atom(elements)127 similars = [128 {'formula': ex.formula_pretty,129 'prop': self.p_label if ex.dict()[self.prop] else self.n_label130 } for ex in similars131 ]132 examples = pd.DataFrame(similars).drop_duplicates().to_dict(orient="records")133 example_selector = MaxMarginalRelevanceExampleSelector.from_examples(134 examples,135 OpenAIEmbeddings(),136 FAISS,137 k=self.k,138 )139 140 prefix=(141 f'You are a bot who can predict if a material is {self.prop_name}.\n'142 f'Given this list of known materials and the information if they are {self.p_label} or {self.n_label}, \n'143 f'you need to answer the question if the last material is {self.prop_name}:'144 )145 prompt_template=PromptTemplate(146 input_variables=["formula", "prop"],147 template=f"Is {{formula}} a {self.prop_name} material?@@@\n{{prop}}###",148 )149 suffix = f"Is {{formula}} a {self.prop_name} material?@@@\n"150 prompt = FewShotPromptTemplate(151 # examples=examples,152 example_prompt=prompt_template,153 example_selector=example_selector,154 prefix=prefix,155 suffix=suffix,156 input_variables=["formula"])157 158 return prompt.format(formula=formula)159 160class MAPI_reg_tools(MAPITools):161 # TODO: deal with units162 def __init__(self, prop, prop_name):163 super().__init__()164 self.prop = prop165 self.prop_name = prop_name166 167 def check_prop_by_formula(self, formula):168 ''' This functions searches in the material project's API for the formula and returns if it is stable or not'''169 with MPRester(os.getenv("MAPI_API_KEY")) as mpr:170 docs = mpr.summary.search(formula=formula, fields=["formula_pretty", self.prop])171 if docs:172 if docs[0].formula_pretty == formula:173 return docs[0].dict()[self.prop]174 elif docs[0].dict()[self.prop] is None:175 return f"There is no record of {self.prop_name} for {formula}"176 return f"Could not find any material while searching {formula}"177 178 def create_context_prompt(self, formula):179 f'''This function received a material formula as input and create a prompt to be inputed in the LLM_predict tool to predict the {self.prop_name} of the material '''180 elements = self.get_material_atoms(formula)181 similars = self.search_similars_by_atom(elements)182 similars = [183 {'formula': ex.formula_pretty,184 'prop': f"{ex.dict()[self.prop]:2f}" if ex.dict()[self.prop] is not None else None185 } for ex in similars186 ]187 examples = pd.DataFrame(similars).drop_duplicates().dropna().to_dict(orient="records")188 189 example_selector = MaxMarginalRelevanceExampleSelector.from_examples(190 examples,191 OpenAIEmbeddings(),192 FAISS,193 k=self.k,194 )195 196 prefix=(197 f'You are a bot who can predict the {self.prop_name} of a material .\n'198 f'Given this list of known materials and the measurement of their {self.prop_name}, \n'199 f'you need to answer the what is the {self.prop_name} of the material:'200 'The answer should be numeric and finish with ###'201 )202 prompt_template=PromptTemplate(203 input_variables=["formula", "prop"],204 template=f"What is the {self.prop_name} for {{formula}}?@@@\n{{prop}}###",205 )206 suffix = f"What is the {self.prop_name} for {{formula}}?@@@\n"207 prompt = FewShotPromptTemplate(208 # examples=examples,209 example_prompt=prompt_template,210 example_selector=example_selector,211 prefix=prefix,212 suffix=suffix,213 input_variables=["formula"])214 215 return prompt.format(formula=formula)216 