Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
kg.py142 linesDownload Raw Back to memory
1from typing import Any, Dict, List, Type, Union2 3from langchain_core.language_models import BaseLanguageModel4from langchain_core.messages import BaseMessage, SystemMessage, get_buffer_string5from langchain_core.prompts import BasePromptTemplate6from pydantic import Field7 8from langchain_community.graphs import NetworkxEntityGraph9from langchain_community.graphs.networkx_graph import (10    KnowledgeTriple,11    get_entities,12    parse_triples,13)14 15try:16    from langchain_classic.chains.llm import LLMChain17    from langchain_classic.memory.chat_memory import BaseChatMemory18    from langchain_classic.memory.prompt import (19        ENTITY_EXTRACTION_PROMPT,20        KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT,21    )22    from langchain_classic.memory.utils import get_prompt_input_key23 24    class ConversationKGMemory(BaseChatMemory):25        """Knowledge graph conversation memory.26 27        Integrates with external knowledge graph to store and retrieve28        information about knowledge triples in the conversation.29        """30 31        k: int = 232        human_prefix: str = "Human"33        ai_prefix: str = "AI"34        kg: NetworkxEntityGraph = Field(default_factory=NetworkxEntityGraph)35        knowledge_extraction_prompt: BasePromptTemplate = (36            KNOWLEDGE_TRIPLE_EXTRACTION_PROMPT37        )38        entity_extraction_prompt: BasePromptTemplate = ENTITY_EXTRACTION_PROMPT39        llm: BaseLanguageModel40        summary_message_cls: Type[BaseMessage] = SystemMessage41        """Number of previous utterances to include in the context."""42        memory_key: str = "history"  #: :meta private:43 44        def load_memory_variables(self, inputs: Dict[str, Any]) -> Dict[str, Any]:45            """Return history buffer."""46            entities = self._get_current_entities(inputs)47 48            summary_strings = []49            for entity in entities:50                knowledge = self.kg.get_entity_knowledge(entity)51                if knowledge:52                    summary = f"On {entity}: {'. '.join(knowledge)}."53                    summary_strings.append(summary)54            context: Union[str, List]55            if not summary_strings:56                context = [] if self.return_messages else ""57            elif self.return_messages:58                context = [59                    self.summary_message_cls(content=text) for text in summary_strings60                ]61            else:62                context = "\n".join(summary_strings)63 64            return {self.memory_key: context}65 66        @property67        def memory_variables(self) -> List[str]:68            """Will always return list of memory variables.69 70            :meta private:71            """72            return [self.memory_key]73 74        def _get_prompt_input_key(self, inputs: Dict[str, Any]) -> str:75            """Get the input key for the prompt."""76            if self.input_key is None:77                return get_prompt_input_key(inputs, self.memory_variables)78            return self.input_key79 80        def _get_prompt_output_key(self, outputs: Dict[str, Any]) -> str:81            """Get the output key for the prompt."""82            if self.output_key is None:83                if len(outputs) != 1:84                    raise ValueError(f"One output key expected, got {outputs.keys()}")85                return list(outputs.keys())[0]86            return self.output_key87 88        def get_current_entities(self, input_string: str) -> List[str]:89            chain = LLMChain(llm=self.llm, prompt=self.entity_extraction_prompt)90            buffer_string = get_buffer_string(91                self.chat_memory.messages[-self.k * 2 :],92                human_prefix=self.human_prefix,93                ai_prefix=self.ai_prefix,94            )95            output = chain.predict(96                history=buffer_string,97                input=input_string,98            )99            return get_entities(output)100 101        def _get_current_entities(self, inputs: Dict[str, Any]) -> List[str]:102            """Get the current entities in the conversation."""103            prompt_input_key = self._get_prompt_input_key(inputs)104            return self.get_current_entities(inputs[prompt_input_key])105 106        def get_knowledge_triplets(self, input_string: str) -> List[KnowledgeTriple]:107            chain = LLMChain(llm=self.llm, prompt=self.knowledge_extraction_prompt)108            buffer_string = get_buffer_string(109                self.chat_memory.messages[-self.k * 2 :],110                human_prefix=self.human_prefix,111                ai_prefix=self.ai_prefix,112            )113            output = chain.predict(114                history=buffer_string,115                input=input_string,116                verbose=True,117            )118            knowledge = parse_triples(output)119            return knowledge120 121        def _get_and_update_kg(self, inputs: Dict[str, Any]) -> None:122            """Get and update knowledge graph from the conversation history."""123            prompt_input_key = self._get_prompt_input_key(inputs)124            knowledge = self.get_knowledge_triplets(inputs[prompt_input_key])125            for triple in knowledge:126                self.kg.add_triple(triple)127 128        def save_context(self, inputs: Dict[str, Any], outputs: Dict[str, str]) -> None:129            """Save context from this conversation to buffer."""130            super().save_context(inputs, outputs)131            self._get_and_update_kg(inputs)132 133        def clear(self) -> None:134            """Clear memory contents."""135            super().clear()136            self.kg.clear()137 138except ImportError:139    # Placeholder object140    class ConversationKGMemory:  # type: ignore[no-redef]141        pass142 
codekingpro/portable-devtools · Team Ai