mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
feat(core): integrate AgentScope LLM support with enhanced memory management
This commit is contained in:
parent
3dc3c4bf52
commit
57f8a7b42c
30 changed files with 467 additions and 808 deletions
|
|
@ -1,11 +1,28 @@
|
|||
as_llms:
|
||||
default:
|
||||
backend: openai
|
||||
model_name: qwen3.5-plus
|
||||
|
||||
as_llm_formatters:
|
||||
default:
|
||||
backend: openai
|
||||
|
||||
embedding_models:
|
||||
default:
|
||||
backend: openai
|
||||
dimensions: 1024
|
||||
use_dimensions: false
|
||||
enable_cache: true
|
||||
max_batch_size: 10
|
||||
max_cache_size: 2000
|
||||
max_input_length: 8192
|
||||
|
||||
file_stores:
|
||||
default:
|
||||
backend: chroma
|
||||
embedding_model: default
|
||||
store_name: "reme"
|
||||
|
||||
|
||||
file_watchers:
|
||||
default:
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Core"""
|
||||
|
||||
from . import as_llm
|
||||
from . import as_llm_formatter
|
||||
from . import embedding
|
||||
from . import enumeration
|
||||
from . import file_store
|
||||
|
|
@ -21,6 +22,8 @@ from .service_context import ServiceContext
|
|||
|
||||
__all__ = [
|
||||
# Submodules
|
||||
"as_llm",
|
||||
"as_llm_formatter",
|
||||
"embedding",
|
||||
"enumeration",
|
||||
"file_watcher",
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""High-level entry point for configuring and running ReMe services and flows."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
|
||||
|
|
@ -24,24 +25,26 @@ class Application:
|
|||
"""Application wrapper that wires together service context, flows, and runtimes."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*args,
|
||||
llm_api_key: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
working_dir: str | None = None,
|
||||
config_path: str | None = None,
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
parser: type[PydanticConfigParser] | None = None,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_vector_store_config: dict | None = None,
|
||||
default_file_store_config: dict | None = None,
|
||||
default_token_counter_config: dict | None = None,
|
||||
default_file_watcher_config: dict | None = None,
|
||||
**kwargs,
|
||||
self,
|
||||
*args,
|
||||
llm_api_key: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
working_dir: str | None = None,
|
||||
config_path: str | None = None,
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
parser: type[PydanticConfigParser] | None = None,
|
||||
default_as_llm_config: dict | None = None,
|
||||
default_as_llm_formatter_config: dict | None = None,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_vector_store_config: dict | None = None,
|
||||
default_file_store_config: dict | None = None,
|
||||
default_token_counter_config: dict | None = None,
|
||||
default_file_watcher_config: dict | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.service_context = ServiceContext(
|
||||
*args,
|
||||
|
|
@ -55,6 +58,8 @@ class Application:
|
|||
config_path=config_path,
|
||||
enable_logo=enable_logo,
|
||||
log_to_console=log_to_console,
|
||||
default_as_llm_config=default_as_llm_config,
|
||||
default_as_llm_formatter_config=default_as_llm_formatter_config,
|
||||
default_llm_config=default_llm_config,
|
||||
default_embedding_model_config=default_embedding_model_config,
|
||||
default_vector_store_config=default_vector_store_config,
|
||||
|
|
@ -137,8 +142,8 @@ class Application:
|
|||
ray.init(num_cpus=self.service_config.ray_max_workers)
|
||||
|
||||
if (
|
||||
self.service_context.thread_pool is None
|
||||
or self.service_context.thread_pool._shutdown # pylint: disable=protected-access
|
||||
self.service_context.thread_pool is None
|
||||
or self.service_context.thread_pool._shutdown # pylint: disable=protected-access
|
||||
):
|
||||
self.service_context.thread_pool = ThreadPoolExecutor(
|
||||
max_workers=self.service_config.thread_pool_max_workers,
|
||||
|
|
@ -147,6 +152,26 @@ class Application:
|
|||
if self.service_context.service_config.enable_logo:
|
||||
print_logo(service_config=self.service_config)
|
||||
|
||||
for name, config in self.service_config.as_llms.items():
|
||||
if config.backend not in R.as_llms:
|
||||
logger.warning(f"AS LLM backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
if not config_dict.get("api_key", ""):
|
||||
config_dict["api_key"] = os.getenv("LLM_API_KEY", "")
|
||||
if "client_kwargs" not in config_dict:
|
||||
config_dict["client_kwargs"] = {}
|
||||
if not config_dict["client_kwargs"].get("base_url", ""):
|
||||
config_dict["client_kwargs"]["base_url"] = os.getenv("LLM_BASE_URL", "")
|
||||
self.service_context.as_llms[name] = R.as_llms[config.backend](**config_dict)
|
||||
|
||||
for name, config in self.service_config.as_llm_formatters.items():
|
||||
if config.backend not in R.as_llm_formatters:
|
||||
logger.warning(f"AS LLM formatter backend {config.backend} is not supported.")
|
||||
else:
|
||||
config_dict = config.model_dump(exclude={"backend"})
|
||||
self.service_context.as_llm_formatters[name] = R.as_llm_formatters[config.backend](**config_dict)
|
||||
|
||||
for name, config in self.service_config.llms.items():
|
||||
if config.backend not in R.llms:
|
||||
logger.warning(f"LLM backend {config.backend} is not supported.")
|
||||
|
|
@ -294,10 +319,10 @@ class Application:
|
|||
stream_queue = asyncio.Queue()
|
||||
task = asyncio.create_task(flow.call(stream_queue=stream_queue, **kwargs))
|
||||
async for chunk in execute_stream_task(
|
||||
stream_queue=stream_queue,
|
||||
task=task,
|
||||
task_name=name,
|
||||
output_format="str",
|
||||
stream_queue=stream_queue,
|
||||
task=task,
|
||||
task_name=name,
|
||||
output_format="str",
|
||||
):
|
||||
yield chunk
|
||||
|
||||
|
|
|
|||
7
reme/core/as_llm/__init__.py
Normal file
7
reme/core/as_llm/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
from agentscope.model import DashScopeChatModel
|
||||
from agentscope.model import OpenAIChatModel
|
||||
|
||||
from ..registry_factory import R
|
||||
|
||||
R.as_llms.register(OpenAIChatModel, "openai")
|
||||
R.as_llms.register(DashScopeChatModel, "dashscope")
|
||||
7
reme/core/as_llm_formatter/__init__.py
Normal file
7
reme/core/as_llm_formatter/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
from agentscope.formatter import DashScopeChatFormatter
|
||||
from agentscope.formatter import OpenAIChatFormatter
|
||||
|
||||
from ..registry_factory import R
|
||||
|
||||
R.as_llm_formatters.register(OpenAIChatFormatter, "openai")
|
||||
R.as_llm_formatters.register(DashScopeChatFormatter, "dashscope")
|
||||
|
|
@ -21,7 +21,8 @@ from ..service_context import ServiceContext
|
|||
from ..token_counter import BaseTokenCounter
|
||||
from ..utils import camel_to_snake, CacheHandler, timer
|
||||
from ..vector_store import BaseVectorStore
|
||||
|
||||
from agentscope.model import ChatModelBase
|
||||
from agentscope.formatter import FormatterBase
|
||||
|
||||
class BaseOp(metaclass=ABCMeta):
|
||||
"""Base operator class for LLM workflow execution and composition."""
|
||||
|
|
@ -42,6 +43,8 @@ class BaseOp(metaclass=ABCMeta):
|
|||
language: str = "",
|
||||
prompt_name: str = "",
|
||||
prompt_path: str = "",
|
||||
as_llm: str | ChatModelBase = "default",
|
||||
as_llm_formatter: str | FormatterBase = "default",
|
||||
llm: str | BaseLLM = "default",
|
||||
embedding_model: str | BaseEmbeddingModel = "default",
|
||||
vector_store: str | BaseVectorStore = "default",
|
||||
|
|
@ -64,6 +67,8 @@ class BaseOp(metaclass=ABCMeta):
|
|||
self.language = language
|
||||
self.prompt = self._get_prompt_handler(prompt_name, prompt_path)
|
||||
|
||||
self._as_llm = as_llm
|
||||
self._as_llm_formatter = as_llm_formatter
|
||||
self._llm = llm
|
||||
self._embedding_model = embedding_model
|
||||
self._vector_store = vector_store
|
||||
|
|
@ -129,6 +134,20 @@ class BaseOp(metaclass=ABCMeta):
|
|||
"""Access the service configuration."""
|
||||
return self.service_context.service_config
|
||||
|
||||
@property
|
||||
def as_llm(self) -> ChatModelBase:
|
||||
"""Get the AgentScope LLM instance from ServiceContext."""
|
||||
if isinstance(self._as_llm, str):
|
||||
self._as_llm = self.service_context.as_llms[self._as_llm]
|
||||
return self._as_llm
|
||||
|
||||
@property
|
||||
def as_llm_formatter(self) -> FormatterBase:
|
||||
"""Get the AgentScope LLM formatter instance from ServiceContext."""
|
||||
if isinstance(self._as_llm_formatter, str):
|
||||
self._as_llm_formatter = self.service_context.as_llm_formatters[self._as_llm_formatter]
|
||||
return self._as_llm_formatter
|
||||
|
||||
@property
|
||||
def llm(self) -> BaseLLM:
|
||||
"""Get the LLM instance from ServiceContext."""
|
||||
|
|
|
|||
|
|
@ -34,6 +34,8 @@ class RegistryFactory:
|
|||
|
||||
def __init__(self):
|
||||
self.llms = Registry()
|
||||
self.as_llms = Registry()
|
||||
self.as_llm_formatters = Registry()
|
||||
self.embedding_models = Registry()
|
||||
self.vector_stores = Registry()
|
||||
self.file_stores = Registry()
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""schema"""
|
||||
|
||||
from .as_msg_stat import AsBlockStat, AsMsgStat
|
||||
from .cut_point_result import CutPointResult
|
||||
from .file_metadata import FileMetadata
|
||||
from .memory_chunk import MemoryChunk
|
||||
|
|
@ -27,6 +28,8 @@ from .truncation_result import TruncationResult
|
|||
from .vector_node import VectorNode
|
||||
|
||||
__all__ = [
|
||||
"AsBlockStat",
|
||||
"AsMsgStat",
|
||||
"CutPointResult",
|
||||
"CmdConfig",
|
||||
"ContentBlock",
|
||||
|
|
|
|||
|
|
@ -3,24 +3,6 @@ from pydantic import BaseModel, Field
|
|||
_DEFAULT_MAX_BLOCK_TEXT_PREVIEW_LENGTH = 100
|
||||
_DEFAULT_MAX_FORMATTER_TEXT_LENGTH = 2000
|
||||
|
||||
# Unique marker for truncated text
|
||||
TRUNCATION_MARKER_START = "<<<TRUNCATED>>>"
|
||||
TRUNCATION_MARKER_END = "<<<END_TRUNCATED>>>"
|
||||
|
||||
|
||||
def _truncate_text(text: str, max_length: int) -> str:
|
||||
"""Truncate text to max length, keeping head and tail portions."""
|
||||
text = str(text) if text else ""
|
||||
if not text or len(text) <= max_length:
|
||||
return text
|
||||
half_length = max_length // 2
|
||||
truncated_chars = len(text) - max_length
|
||||
return (
|
||||
f"{text[:half_length]}\n\n{TRUNCATION_MARKER_START} "
|
||||
f"({truncated_chars} characters omitted) "
|
||||
f"{TRUNCATION_MARKER_END}\n\n{text[-half_length:]}"
|
||||
)
|
||||
|
||||
|
||||
class AsBlockStat(BaseModel):
|
||||
block_type: str = Field(default=...)
|
||||
|
|
@ -41,18 +23,20 @@ class AsBlockStat(BaseModel):
|
|||
|
||||
def format(self, max_length: int = _DEFAULT_MAX_FORMATTER_TEXT_LENGTH, include_thinking: bool = True) -> str:
|
||||
"""Format block content to string representation."""
|
||||
from ..utils import truncate_text
|
||||
|
||||
if self.block_type == "text":
|
||||
return _truncate_text(self.text, max_length) if self.text else ""
|
||||
return truncate_text(self.text, max_length) if self.text else ""
|
||||
if self.block_type == "thinking":
|
||||
if include_thinking and self.text:
|
||||
return f"<thinking>\n{_truncate_text(self.text, max_length)}\n</thinking>"
|
||||
return f"<thinking>\n{truncate_text(self.text, max_length)}\n</thinking>"
|
||||
return ""
|
||||
if self.block_type in ("image", "audio", "video"):
|
||||
return f"[{self.block_type}] {self.media_url}" if self.media_url else f"[{self.block_type}]"
|
||||
if self.block_type == "tool_use":
|
||||
return f" - tool_call={self.tool_name} params={_truncate_text(self.tool_input, max_length)}"
|
||||
return f" - tool_call={self.tool_name} params={truncate_text(self.tool_input, max_length)}"
|
||||
if self.block_type == "tool_result":
|
||||
output = _truncate_text(self.tool_output, max_length)
|
||||
output = truncate_text(self.tool_output, max_length)
|
||||
return f" - tool_result={self.tool_name} output={output}" if output else ""
|
||||
return ""
|
||||
|
||||
|
|
|
|||
|
|
@ -58,69 +58,60 @@ class FlowConfig(ToolCall):
|
|||
cache_expire_hours: float = Field(default=0.1)
|
||||
|
||||
|
||||
class LLMConfig(BaseModel):
|
||||
class BasicConfig(BaseModel):
|
||||
"""Configuration for basic service settings and parameters."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="")
|
||||
|
||||
|
||||
class ModelConfig(BasicConfig):
|
||||
"""Configuration for model-based services with backend and model name."""
|
||||
|
||||
model_name: str = Field(default="")
|
||||
|
||||
|
||||
class LLMConfig(ModelConfig):
|
||||
"""Configuration for Large Language Model backend and model identification."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="")
|
||||
model_name: str = Field(default="")
|
||||
|
||||
|
||||
class EmbeddingModelConfig(BaseModel):
|
||||
class EmbeddingModelConfig(ModelConfig):
|
||||
"""Configuration for embedding model backends and identity."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="")
|
||||
model_name: str = Field(default="")
|
||||
|
||||
|
||||
class VectorStoreConfig(BaseModel):
|
||||
"""Configuration for vector database storage and associated embeddings."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="local")
|
||||
collection_name: str = Field(default="reme")
|
||||
embedding_model: str = Field(default="default")
|
||||
|
||||
|
||||
class FileStoreConfig(BaseModel):
|
||||
"""Configuration for file store database storage and associated embeddings."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="sqlite")
|
||||
store_name: str = Field(default="reme")
|
||||
embedding_model: str = Field(default="default")
|
||||
|
||||
|
||||
class TokenCounterConfig(BaseModel):
|
||||
class TokenCounterConfig(ModelConfig):
|
||||
"""Configuration for token counting services and model mapping."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="base")
|
||||
model_name: str = Field(default="")
|
||||
class StoreConfig(BasicConfig):
|
||||
"""Configuration for storage services with embedding model support."""
|
||||
|
||||
embedding_model: str = Field(default="default")
|
||||
|
||||
|
||||
class FileWatcherConfig(BaseModel):
|
||||
class VectorStoreConfig(StoreConfig):
|
||||
"""Configuration for vector database storage and associated embeddings."""
|
||||
|
||||
collection_name: str = Field(default="reme")
|
||||
|
||||
|
||||
class FileStoreConfig(StoreConfig):
|
||||
"""Configuration for file store database storage and associated embeddings."""
|
||||
|
||||
store_name: str = Field(default="reme")
|
||||
|
||||
|
||||
class FileWatcherConfig(BasicConfig):
|
||||
"""Configuration for file watcher service."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="")
|
||||
file_store: str = Field(default="")
|
||||
watch_paths: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ServiceConfig(BaseModel):
|
||||
class ServiceConfig(BasicConfig):
|
||||
"""Root configuration schema aggregating all service-level settings and components."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
backend: str = Field(default="")
|
||||
app_name: str = Field(default=os.getenv("APP_NAME", "ReMe"))
|
||||
working_dir: str = Field(default=".reme")
|
||||
enable_logo: bool = Field(default=True)
|
||||
|
|
@ -137,6 +128,8 @@ class ServiceConfig(BaseModel):
|
|||
cmd: CmdConfig = Field(default_factory=CmdConfig)
|
||||
ops: dict[str, OpConfig] = Field(default_factory=dict)
|
||||
flows: dict[str, FlowConfig] = Field(default_factory=dict)
|
||||
as_llms: dict[str, BasicConfig] = Field(default_factory=dict)
|
||||
as_llm_formatters: dict[str, BasicConfig] = Field(default_factory=dict)
|
||||
llms: dict[str, LLMConfig] = Field(default_factory=dict)
|
||||
embedding_models: dict[str, EmbeddingModelConfig] = Field(default_factory=dict)
|
||||
vector_stores: dict[str, VectorStoreConfig] = Field(default_factory=dict)
|
||||
|
|
|
|||
|
|
@ -11,6 +11,8 @@ from .schema import ServiceConfig
|
|||
from .utils import load_env, PydanticConfigParser
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agentscope.model import ChatModelBase
|
||||
from agentscope.formatter import FormatterBase
|
||||
from .llm import BaseLLM
|
||||
from .embedding import BaseEmbeddingModel
|
||||
from .vector_store import BaseVectorStore
|
||||
|
|
@ -36,6 +38,8 @@ class ServiceContext(BaseDict):
|
|||
config_path: str | None = None,
|
||||
enable_logo: bool = True,
|
||||
log_to_console: bool = True,
|
||||
default_as_llm_config: dict | None = None,
|
||||
default_as_llm_formatter_config: dict | None = None,
|
||||
default_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_vector_store_config: dict | None = None,
|
||||
|
|
@ -64,6 +68,10 @@ class ServiceContext(BaseDict):
|
|||
if args:
|
||||
input_args.extend(args)
|
||||
|
||||
if default_as_llm_config:
|
||||
self._update_section_config(kwargs, "as_llms", **default_as_llm_config)
|
||||
if default_as_llm_formatter_config:
|
||||
self._update_section_config(kwargs, "as_llm_formatters", **default_as_llm_formatter_config)
|
||||
if default_llm_config:
|
||||
self._update_section_config(kwargs, "llms", **default_llm_config)
|
||||
if default_embedding_model_config:
|
||||
|
|
@ -90,6 +98,8 @@ class ServiceContext(BaseDict):
|
|||
self.service_config: ServiceConfig = service_config
|
||||
|
||||
self.thread_pool: ThreadPoolExecutor | None = None
|
||||
self.as_llms: dict[str, "ChatModelBase"] = {}
|
||||
self.as_llm_formatters: dict[str, "FormatterBase"] = {}
|
||||
self.llms: dict[str, "BaseLLM"] = {}
|
||||
self.embedding_models: dict[str, "BaseEmbeddingModel"] = {}
|
||||
self.token_counters: dict[str, "BaseTokenCounter"] = {}
|
||||
|
|
|
|||
|
|
@ -11,12 +11,15 @@ from .horse import play_horse_easter_egg
|
|||
from .http_client import HttpClient
|
||||
from .llm_utils import extract_content, format_messages, deduplicate_memories
|
||||
from .logger_utils import init_logger
|
||||
from .std_logger import get_logger as get_std_logger
|
||||
from .logo_utils import print_logo
|
||||
from .mcp_client import MCPClient
|
||||
from .pydantic_config_parser import PydanticConfigParser
|
||||
from .pydantic_utils import create_pydantic_model
|
||||
from .singleton import singleton
|
||||
from .time import timer, get_now_time
|
||||
from .hf_token_counter_utils import get_hf_token_counter
|
||||
from .truncate_text_utils import truncate_text, is_truncated
|
||||
|
||||
__all__ = [
|
||||
"convert_dashscope_to_agentscope",
|
||||
|
|
@ -39,6 +42,7 @@ __all__ = [
|
|||
"format_messages",
|
||||
"deduplicate_memories",
|
||||
"init_logger",
|
||||
"get_std_logger",
|
||||
"print_logo",
|
||||
"MCPClient",
|
||||
"PydanticConfigParser",
|
||||
|
|
@ -46,4 +50,7 @@ __all__ = [
|
|||
"singleton",
|
||||
"timer",
|
||||
"get_now_time",
|
||||
"get_hf_token_counter",
|
||||
"truncate_text",
|
||||
"is_truncated",
|
||||
]
|
||||
|
|
|
|||
23
reme/core/utils/hf_token_counter_utils.py
Normal file
23
reme/core/utils/hf_token_counter_utils.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
"""Utility functions for working with text."""
|
||||
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
_token_counter = None
|
||||
|
||||
|
||||
def get_hf_token_counter(
|
||||
pretrained_model_name_or_path="Qwen/Qwen2.5-7B-Instruct",
|
||||
use_mirror=True,
|
||||
use_fast=True,
|
||||
trust_remote_code=True,
|
||||
):
|
||||
"""Get or initialize the global token counter instance."""
|
||||
global _token_counter
|
||||
if _token_counter is None:
|
||||
_token_counter = HuggingFaceTokenCounter(
|
||||
pretrained_model_name_or_path=pretrained_model_name_or_path,
|
||||
use_mirror=use_mirror,
|
||||
use_fast=use_fast,
|
||||
trust_remote_code=trust_remote_code,
|
||||
)
|
||||
return _token_counter
|
||||
109
reme/core/utils/std_logger.py
Normal file
109
reme/core/utils/std_logger.py
Normal file
|
|
@ -0,0 +1,109 @@
|
|||
"""Standard logging module configuration with loguru-like features."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from logging.handlers import TimedRotatingFileHandler
|
||||
|
||||
# Store created logger instances
|
||||
_loggers: dict[str, logging.Logger] = {}
|
||||
|
||||
|
||||
class CustomFormatter(logging.Formatter):
|
||||
"""Custom formatter with colorized output support."""
|
||||
|
||||
# ANSI color codes
|
||||
COLORS = {
|
||||
logging.DEBUG: "\033[36m", # Cyan
|
||||
logging.INFO: "\033[32m", # Green
|
||||
logging.WARNING: "\033[33m", # Yellow
|
||||
logging.ERROR: "\033[31m", # Red
|
||||
logging.CRITICAL: "\033[35m", # Magenta
|
||||
}
|
||||
RESET = "\033[0m"
|
||||
|
||||
def __init__(self, fmt: str, colorize: bool = False):
|
||||
super().__init__(fmt)
|
||||
self.colorize = colorize
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
# Add custom attribute: simplified filename and line number
|
||||
record.file_line = f"{record.filename}:{record.lineno}"
|
||||
|
||||
if self.colorize:
|
||||
color = self.COLORS.get(record.levelno, self.RESET)
|
||||
record.levelname = f"{color}{record.levelname}{self.RESET}"
|
||||
|
||||
return super().format(record)
|
||||
|
||||
|
||||
def get_logger(
|
||||
name: str = "reme",
|
||||
log_dir: str = "logs",
|
||||
level: str = "INFO",
|
||||
log_to_console: bool = True,
|
||||
log_to_file: bool = True,
|
||||
log_file_prefix: str = "reme",
|
||||
rotation: str = "midnight",
|
||||
retention_days: int = 7,
|
||||
) -> logging.Logger:
|
||||
"""Get a configured logger instance.
|
||||
|
||||
Args:
|
||||
name: Logger name for distinguishing different loggers.
|
||||
log_dir: Directory path for log files.
|
||||
level: Logging level (DEBUG, INFO, WARNING, ERROR, CRITICAL).
|
||||
log_to_console: Whether to output logs to console.
|
||||
log_to_file: Whether to output logs to file.
|
||||
log_file_prefix: Prefix for log file names (e.g., 'reme' -> 'reme_2024-01-01.log').
|
||||
rotation: Log rotation time, defaults to midnight.
|
||||
retention_days: Number of days to retain log files.
|
||||
|
||||
Returns:
|
||||
Configured Logger instance.
|
||||
"""
|
||||
# Return existing logger if already created
|
||||
if name in _loggers:
|
||||
return _loggers[name]
|
||||
|
||||
# Create new logger without using root logger
|
||||
logger = logging.getLogger(name)
|
||||
logger.setLevel(getattr(logging, level.upper(), logging.INFO))
|
||||
logger.propagate = False # Do not propagate to root logger
|
||||
|
||||
# Clear existing handlers
|
||||
logger.handlers.clear()
|
||||
|
||||
# Log format
|
||||
log_format = "%(asctime)s | %(levelname)s | %(file_line)s | %(funcName)s | %(message)s"
|
||||
|
||||
# Configure file logging
|
||||
if log_to_file:
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
log_filename = f"{log_file_prefix}_{current_ts}.log"
|
||||
log_filepath = os.path.join(log_dir, log_filename)
|
||||
|
||||
file_handler = TimedRotatingFileHandler(
|
||||
log_filepath,
|
||||
when=rotation,
|
||||
interval=1,
|
||||
backupCount=retention_days,
|
||||
encoding="utf-8",
|
||||
)
|
||||
file_handler.setLevel(getattr(logging, level.upper(), logging.INFO))
|
||||
file_handler.setFormatter(CustomFormatter(log_format, colorize=False))
|
||||
file_handler.suffix = "%Y-%m-%d"
|
||||
logger.addHandler(file_handler)
|
||||
|
||||
# Configure console logging
|
||||
if log_to_console:
|
||||
console_handler = logging.StreamHandler(sys.stdout)
|
||||
console_handler.setLevel(getattr(logging, level.upper(), logging.INFO))
|
||||
console_handler.setFormatter(CustomFormatter(log_format, colorize=True))
|
||||
logger.addHandler(console_handler)
|
||||
|
||||
# Cache logger
|
||||
_loggers[name] = logger
|
||||
return logger
|
||||
53
reme/core/utils/truncate_text_utils.py
Normal file
53
reme/core/utils/truncate_text_utils.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
from .std_logger import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
TRUNCATION_MARKER_START = "<<<TRUNCATED>>>"
|
||||
TRUNCATION_MARKER_END = "<<<END_TRUNCATED>>>"
|
||||
|
||||
|
||||
def truncate_text(text: str, max_length: int) -> str:
|
||||
"""Truncate text to max length, keeping head and tail portions.
|
||||
|
||||
Args:
|
||||
text: The text to truncate
|
||||
max_length: Maximum allowed length
|
||||
|
||||
Returns:
|
||||
Truncated text with unique markers indicating truncation
|
||||
"""
|
||||
text = str(text) if text else ""
|
||||
if not text:
|
||||
return text
|
||||
|
||||
if len(text) <= max_length:
|
||||
return text
|
||||
|
||||
half_length = max_length // 2
|
||||
truncated_chars = len(text) - max_length
|
||||
logger.debug(
|
||||
"Text truncated: original %d chars, kept head %d + tail %d, removed %d chars.",
|
||||
len(text),
|
||||
half_length,
|
||||
half_length,
|
||||
truncated_chars,
|
||||
)
|
||||
return (
|
||||
f"{text[:half_length]}\n\n{TRUNCATION_MARKER_START} "
|
||||
f"({truncated_chars} characters omitted) "
|
||||
f"{TRUNCATION_MARKER_END}\n\n{text[-half_length:]}"
|
||||
)
|
||||
|
||||
|
||||
def is_truncated(text: str) -> bool:
|
||||
"""Check if the text has been truncated (contains truncation markers).
|
||||
|
||||
Args:
|
||||
text: The text to check
|
||||
|
||||
Returns:
|
||||
bool: True if text contains truncation markers, False otherwise
|
||||
"""
|
||||
if not text:
|
||||
return False
|
||||
return TRUNCATION_MARKER_START in text and TRUNCATION_MARKER_END in text
|
||||
|
|
@ -5,22 +5,17 @@ including memory formatting, compaction, summarization, and file I/O operations.
|
|||
|
||||
Components:
|
||||
- ReMeInMemoryMemory: Extended InMemoryMemory with bugfixes and summary support
|
||||
- ReMeOpenAIChatFormatter: Converts message lists to formatted strings with token limiting
|
||||
- AsMsgHandler: Handles AgentScope message statistics, formatting, and context checking
|
||||
- Summarizer: Generates memory summaries using LLM
|
||||
- Compactor: Compacts memory content to reduce token usage
|
||||
- ToolResultCompactor: Truncates large tool results and saves full content to files
|
||||
- FileIO: File I/O operations with configurable working directory
|
||||
"""
|
||||
|
||||
from . import utils
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
from .compactor import Compactor
|
||||
from .file_io import FileIO
|
||||
from .reme_chat_formatter import ReMeOpenAIChatFormatter
|
||||
from .reme_in_memory_memory import ReMeInMemoryMemory
|
||||
from .summarizer import Summarizer
|
||||
from .tool_result_compactor import ToolResultCompactor
|
||||
from .sub_agent.compactor import Compactor
|
||||
from .sub_agent.summarizer import Summarizer
|
||||
from .sub_agent.tool_result_compactor import ToolResultCompactor
|
||||
|
||||
__all__ = [
|
||||
"AsMsgHandler",
|
||||
|
|
@ -28,7 +23,4 @@ __all__ = [
|
|||
"Summarizer",
|
||||
"Compactor",
|
||||
"ToolResultCompactor",
|
||||
"FileIO",
|
||||
"utils",
|
||||
"ReMeOpenAIChatFormatter",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
import json
|
||||
import logging
|
||||
|
||||
from agentscope.message import Msg
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from ...core.schema.as_msg_stat import AsMsgStat, AsBlockStat
|
||||
from ...core.schema import AsMsgStat, AsBlockStat
|
||||
from ...core.utils import get_std_logger
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger = get_std_logger()
|
||||
|
||||
|
||||
class AsMsgHandler:
|
||||
|
|
@ -179,10 +179,10 @@ class AsMsgHandler:
|
|||
)
|
||||
|
||||
def format_msgs_to_str(
|
||||
self,
|
||||
messages: list[Msg],
|
||||
memory_compact_threshold: int,
|
||||
include_thinking: bool = False,
|
||||
self,
|
||||
messages: list[Msg],
|
||||
memory_compact_threshold: int,
|
||||
include_thinking: bool = False,
|
||||
) -> str:
|
||||
"""Format list of messages to a single formatted string.
|
||||
|
||||
|
|
@ -348,4 +348,4 @@ class AsMsgHandler:
|
|||
accumulated_tokens,
|
||||
)
|
||||
|
||||
return messages_to_compact, messages_to_keep
|
||||
return messages_to_compact, messages_to_keep
|
||||
|
|
|
|||
|
|
@ -1,29 +0,0 @@
|
|||
"""ReMe chat formatter."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from agentscope.formatter import OpenAIChatFormatter
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from .utils import _extract_text_from_messages
|
||||
|
||||
|
||||
class ReMeOpenAIChatFormatter(OpenAIChatFormatter):
|
||||
"""ReMe chat formatter class."""
|
||||
|
||||
async def _count(self, msgs: list[dict[str, Any]]) -> int | None:
|
||||
"""Count the number of tokens in the input messages. If token counter
|
||||
is not provided, `None` will be returned.
|
||||
|
||||
Args:
|
||||
msgs (`list[Msg]`):
|
||||
The input messages to count tokens for.
|
||||
"""
|
||||
if self.token_counter is None:
|
||||
return None
|
||||
|
||||
assert isinstance(self.token_counter, HuggingFaceTokenCounter)
|
||||
text = _extract_text_from_messages(msgs)
|
||||
token_ids = self.token_counter.tokenizer.encode(text)
|
||||
token_count = len(token_ids)
|
||||
return token_count
|
||||
|
|
@ -1,34 +1,30 @@
|
|||
"""Custom memory implementation with bugfixes and extensions."""
|
||||
|
||||
import logging
|
||||
|
||||
from agentscope.agent._react_agent import _MemoryMark
|
||||
from agentscope.agent._react_agent import _MemoryMark # noqa
|
||||
from agentscope.memory import InMemoryMemory
|
||||
from agentscope.message import Msg
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
from ...core.utils import get_std_logger
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger = get_std_logger()
|
||||
|
||||
|
||||
class ReMeInMemoryMemory(InMemoryMemory):
|
||||
"""Extended InMemoryMemory with bugfixes and summary support."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
):
|
||||
def __init__(self, token_counter: HuggingFaceTokenCounter):
|
||||
super().__init__()
|
||||
self._token_counter: HuggingFaceTokenCounter = token_counter
|
||||
self._msg_handler: AsMsgHandler = AsMsgHandler(token_counter)
|
||||
|
||||
async def get_memory(
|
||||
self,
|
||||
mark: str | None = None,
|
||||
exclude_mark: str | None = _MemoryMark.COMPRESSED,
|
||||
prepend_summary: bool = True,
|
||||
**_kwargs,
|
||||
self,
|
||||
mark: str | None = None,
|
||||
exclude_mark: str | None = _MemoryMark.COMPRESSED,
|
||||
prepend_summary: bool = True,
|
||||
**_kwargs,
|
||||
) -> list[Msg]:
|
||||
"""Get the messages from the memory by mark (if provided).
|
||||
|
||||
|
|
@ -192,10 +188,10 @@ Use it as context to maintain continuity.
|
|||
)
|
||||
|
||||
return (
|
||||
f"**Conversation History**\n\n"
|
||||
f"- Total messages: {stats['total_messages']}\n"
|
||||
f"- Estimated tokens: {stats['estimated_tokens']}\n"
|
||||
f"- Max input length: {stats['max_input_length']}\n"
|
||||
f"- Context usage: {stats['context_usage_ratio']:.1f}%\n"
|
||||
f"- Compressed summary tokens: {stats['compressed_summary_tokens']}\n\n" + "\n\n".join(lines)
|
||||
f"**Conversation History**\n\n"
|
||||
f"- Total messages: {stats['total_messages']}\n"
|
||||
f"- Estimated tokens: {stats['estimated_tokens']}\n"
|
||||
f"- Max input length: {stats['max_input_length']}\n"
|
||||
f"- Context usage: {stats['context_usage_ratio']:.1f}%\n"
|
||||
f"- Compressed summary tokens: {stats['compressed_summary_tokens']}\n\n" + "\n\n".join(lines)
|
||||
)
|
||||
|
|
|
|||
0
reme/memory/file_based/sub_agent/__init__.py
Normal file
0
reme/memory/file_based/sub_agent/__init__.py
Normal file
|
|
@ -1,35 +1,28 @@
|
|||
"""Compactor module for memory compaction operations."""
|
||||
|
||||
import logging
|
||||
|
||||
from agentscope.agent import ReActAgent
|
||||
from agentscope.formatter import FormatterBase
|
||||
from agentscope.message import Msg
|
||||
from agentscope.model import ChatModelBase
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
from ...core.op import BaseOp
|
||||
from ..as_msg_handler import AsMsgHandler
|
||||
from ....core.op import BaseOp
|
||||
from ....core.utils import get_std_logger
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger = get_std_logger()
|
||||
|
||||
|
||||
class Compactor(BaseOp):
|
||||
"""Compactor class for compacting memory messages."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
memory_compact_threshold: int,
|
||||
chat_model: ChatModelBase,
|
||||
formatter: FormatterBase,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
**kwargs,
|
||||
self,
|
||||
memory_compact_threshold: int,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.memory_compact_threshold: int = memory_compact_threshold
|
||||
|
||||
self.chat_model: ChatModelBase = chat_model
|
||||
self.formatter: FormatterBase = formatter
|
||||
self.msg_handler = AsMsgHandler(token_counter=token_counter)
|
||||
|
||||
async def execute(self):
|
||||
|
|
@ -50,9 +43,9 @@ class Compactor(BaseOp):
|
|||
|
||||
agent = ReActAgent(
|
||||
name="reme_compactor",
|
||||
model=self.chat_model,
|
||||
model=self.as_llm,
|
||||
sys_prompt=self.get_prompt("system_prompt"),
|
||||
formatter=self.formatter,
|
||||
formatter=self.as_llm_formatter,
|
||||
)
|
||||
|
||||
if previous_summary:
|
||||
|
|
@ -66,7 +59,7 @@ class Compactor(BaseOp):
|
|||
)
|
||||
else:
|
||||
user_message: str = f"<conversation>\n{history_formatted_str}\n</conversation>\n\n" \
|
||||
+ self.get_prompt("initial_user_message")
|
||||
+ self.get_prompt("initial_user_message")
|
||||
logger.info(f"Compactor sys_prompt={agent.sys_prompt} user_message={user_message}")
|
||||
|
||||
compact_msg: Msg = await agent.reply(
|
||||
|
|
@ -1,42 +1,36 @@
|
|||
"""Summarizer module for memory summarization operations."""
|
||||
|
||||
import datetime
|
||||
import logging
|
||||
|
||||
from agentscope.agent import ReActAgent
|
||||
from agentscope.formatter import FormatterBase
|
||||
from agentscope.message import Msg
|
||||
from agentscope.model import ChatModelBase
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
from agentscope.tool import Toolkit
|
||||
|
||||
from .as_msg_handler import AsMsgHandler
|
||||
from ...core.op import BaseOp
|
||||
from ..as_msg_handler import AsMsgHandler
|
||||
from ....core.op import BaseOp
|
||||
from ....core.utils import get_std_logger
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger = get_std_logger()
|
||||
|
||||
|
||||
class Summarizer(BaseOp):
|
||||
"""Summarizer class for summarizing memory messages."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
working_dir: str,
|
||||
memory_dir: str,
|
||||
memory_compact_threshold: int,
|
||||
chat_model: ChatModelBase,
|
||||
formatter: FormatterBase,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
toolkit: Toolkit,
|
||||
**kwargs,
|
||||
self,
|
||||
working_dir: str,
|
||||
memory_dir: str,
|
||||
memory_compact_threshold: int,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
toolkit: Toolkit,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.working_dir: str = working_dir
|
||||
self.memory_dir: str = memory_dir
|
||||
self.memory_compact_threshold: int = memory_compact_threshold
|
||||
|
||||
self.chat_model: ChatModelBase = chat_model
|
||||
self.formatter: FormatterBase = formatter
|
||||
self.msg_handler = AsMsgHandler(token_counter=token_counter)
|
||||
self.toolkit: Toolkit = toolkit
|
||||
|
||||
|
|
@ -57,9 +51,9 @@ class Summarizer(BaseOp):
|
|||
|
||||
agent = ReActAgent(
|
||||
name="reme_summarizer",
|
||||
model=self.chat_model,
|
||||
model=self.as_llm,
|
||||
sys_prompt="You are a helpful assistant.",
|
||||
formatter=self.formatter,
|
||||
formatter=self.as_llm_formatter,
|
||||
toolkit=self.toolkit,
|
||||
)
|
||||
|
||||
|
|
@ -1,27 +1,27 @@
|
|||
"""Tool Result Compactor: truncate large tool results and save full content to files."""
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
|
||||
from agentscope.message import Msg
|
||||
|
||||
from .utils import is_truncated, truncate_text
|
||||
from ...core.op import BaseOp
|
||||
from ....core.op import BaseOp
|
||||
from ....core.utils import get_std_logger
|
||||
from ....core.utils import truncate_text, is_truncated
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger = get_std_logger()
|
||||
|
||||
|
||||
class ToolResultCompactor(BaseOp):
|
||||
"""Truncate large tool_result outputs and save full content to files."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tool_result_dir: str | Path,
|
||||
tool_result_threshold: int,
|
||||
retention_days: int = 7,
|
||||
**kwargs,
|
||||
self,
|
||||
tool_result_dir: str | Path,
|
||||
tool_result_threshold: int,
|
||||
retention_days: int = 7,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(**kwargs)
|
||||
self.tool_result_dir = Path(tool_result_dir)
|
||||
|
|
@ -1,271 +0,0 @@
|
|||
"""Utility functions for working with text."""
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Unique marker for truncated text
|
||||
TRUNCATION_MARKER_START = "<<<TRUNCATED>>>"
|
||||
TRUNCATION_MARKER_END = "<<<END_TRUNCATED>>>"
|
||||
|
||||
|
||||
def truncate_text(text: str, max_length: int) -> str:
|
||||
"""Truncate text to max length, keeping head and tail portions.
|
||||
|
||||
Args:
|
||||
text: The text to truncate
|
||||
max_length: Maximum allowed length
|
||||
|
||||
Returns:
|
||||
Truncated text with unique markers indicating truncation
|
||||
"""
|
||||
text = str(text) if text else ""
|
||||
if not text:
|
||||
return text
|
||||
|
||||
if len(text) <= max_length:
|
||||
return text
|
||||
|
||||
half_length = max_length // 2
|
||||
truncated_chars = len(text) - max_length
|
||||
logger.debug(
|
||||
"Text truncated: original %d chars, kept head %d + tail %d, removed %d chars.",
|
||||
len(text),
|
||||
half_length,
|
||||
half_length,
|
||||
truncated_chars,
|
||||
)
|
||||
return (
|
||||
f"{text[:half_length]}\n\n{TRUNCATION_MARKER_START} "
|
||||
f"({truncated_chars} characters omitted) "
|
||||
f"{TRUNCATION_MARKER_END}\n\n{text[-half_length:]}"
|
||||
)
|
||||
|
||||
|
||||
def is_truncated(text: str) -> bool:
|
||||
"""Check if the text has been truncated (contains truncation markers).
|
||||
|
||||
Args:
|
||||
text: The text to check
|
||||
|
||||
Returns:
|
||||
bool: True if text contains truncation markers, False otherwise
|
||||
"""
|
||||
if not text:
|
||||
return False
|
||||
return TRUNCATION_MARKER_START in text and TRUNCATION_MARKER_END in text
|
||||
|
||||
|
||||
def _extract_text_from_messages(messages: list[dict]) -> str:
|
||||
"""Extract text content from messages and concatenate into a string.
|
||||
|
||||
Handles various message formats:
|
||||
- Simple string content: {"role": "user", "content": "hello"}
|
||||
- List content with text blocks:
|
||||
{"role": "user", "content": [{"type": "text", "text": "hello"}]}
|
||||
- List content with tool_result blocks:
|
||||
{"role": "user", "content": [{"type": "tool_result", "output": "..."}]}
|
||||
|
||||
Args:
|
||||
messages: List of message dictionaries in chat format.
|
||||
|
||||
Returns:
|
||||
str: Concatenated text content from all messages.
|
||||
"""
|
||||
parts = []
|
||||
for msg in messages:
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
parts.append(content)
|
||||
elif isinstance(content, list):
|
||||
for block in content:
|
||||
if isinstance(block, dict):
|
||||
block_type = block.get("type", "")
|
||||
if block_type == "tool_result":
|
||||
output = block.get("output", "")
|
||||
if isinstance(output, str) and output:
|
||||
parts.append(output)
|
||||
elif isinstance(output, list):
|
||||
for sub in output:
|
||||
if isinstance(sub, dict):
|
||||
sub_text = sub.get("text") or sub.get("content", "")
|
||||
if sub_text:
|
||||
parts.append(str(sub_text))
|
||||
else:
|
||||
text = block.get("text") or block.get("content", "")
|
||||
if text:
|
||||
parts.append(str(text))
|
||||
elif isinstance(block, str):
|
||||
parts.append(block)
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def safe_count_message_tokens(
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
messages: list[dict],
|
||||
) -> int:
|
||||
"""Safely count tokens in messages with fallback estimation.
|
||||
|
||||
This is a wrapper around count_message_tokens that catches exceptions
|
||||
and falls back to a character-based estimation (len // 4) if the
|
||||
tokenizer fails.
|
||||
|
||||
Args:
|
||||
token_counter: Token counter instance.
|
||||
messages: List of message dictionaries in chat format.
|
||||
|
||||
Returns:
|
||||
int: The estimated number of tokens in the messages.
|
||||
"""
|
||||
try:
|
||||
text = _extract_text_from_messages(messages)
|
||||
token_ids = token_counter.tokenizer.encode(text)
|
||||
token_count = len(token_ids)
|
||||
return token_count
|
||||
|
||||
except Exception as e:
|
||||
# Fallback to character-based estimation
|
||||
text = _extract_text_from_messages(messages)
|
||||
estimated_tokens = len(text) // 4
|
||||
logger.warning(
|
||||
"Failed to count tokens: %s, using estimated_tokens=%d",
|
||||
e,
|
||||
estimated_tokens,
|
||||
)
|
||||
return estimated_tokens
|
||||
|
||||
|
||||
def safe_count_str_tokens(
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
text: str,
|
||||
) -> int:
|
||||
"""Safely count tokens in a string with fallback estimation.
|
||||
|
||||
Uses the tokenizer to count tokens in the given text. If the tokenizer
|
||||
fails, falls back to a character-based estimation (len // 4).
|
||||
|
||||
Args:
|
||||
token_counter: Token counter instance.
|
||||
text: The string to count tokens for.
|
||||
|
||||
Returns:
|
||||
int: The estimated number of tokens in the string.
|
||||
"""
|
||||
try:
|
||||
token_ids = token_counter.tokenizer.encode(text)
|
||||
token_count = len(token_ids)
|
||||
return token_count
|
||||
except Exception as e:
|
||||
# Fallback to character-based estimation
|
||||
estimated_tokens = len(text) // 4
|
||||
logger.warning(
|
||||
"Failed to count string tokens: %s, using estimated_tokens=%d",
|
||||
e,
|
||||
estimated_tokens,
|
||||
)
|
||||
return estimated_tokens
|
||||
|
||||
|
||||
def _get_block_tokens( # pylint: disable=too-many-return-statements
|
||||
block: dict,
|
||||
block_type: str,
|
||||
token_counter: HuggingFaceTokenCounter,
|
||||
) -> tuple[int, str]:
|
||||
"""Get token count and content string for different block types.
|
||||
|
||||
Args:
|
||||
block: The content block dict
|
||||
block_type: The type of the block
|
||||
|
||||
Returns:
|
||||
Tuple of (token count, content string)
|
||||
"""
|
||||
if block_type == "text":
|
||||
text = block.get("text", "")
|
||||
return (safe_count_str_tokens(token_counter, text), text) if text else (0, "")
|
||||
|
||||
if block_type == "thinking":
|
||||
thinking = block.get("thinking", "")
|
||||
return (safe_count_str_tokens(token_counter, thinking), thinking) if thinking else (0, "")
|
||||
|
||||
if block_type == "tool_use":
|
||||
# Count input dict and raw_input string
|
||||
input_dict = block.get("input", {})
|
||||
raw_input = block.get("raw_input", "")
|
||||
input_str = str(input_dict) if input_dict else ""
|
||||
total = input_str + raw_input
|
||||
return (safe_count_str_tokens(token_counter, total), total) if total else (0, "")
|
||||
|
||||
if block_type == "tool_result":
|
||||
output = block.get("output")
|
||||
if isinstance(output, str):
|
||||
return (safe_count_str_tokens(token_counter, output), output) if output else (0, "")
|
||||
|
||||
if isinstance(output, list):
|
||||
# Recursively count tokens in nested blocks
|
||||
total_tokens = 0
|
||||
total_str = ""
|
||||
for item in output:
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type", "unknown")
|
||||
item_tokens, item_str = _get_block_tokens(item, item_type, token_counter)
|
||||
total_tokens += item_tokens
|
||||
total_str += item_str
|
||||
return total_tokens, total_str
|
||||
return 0, ""
|
||||
|
||||
if block_type in ("image", "audio", "video"):
|
||||
# For media blocks, count the URL or indicate base64 size
|
||||
source = block.get("source", {})
|
||||
if source.get("type") == "url":
|
||||
url = source.get("url", "")
|
||||
return safe_count_str_tokens(token_counter, url), url
|
||||
if source.get("type") == "base64":
|
||||
# Base64 data can be large, return approximate token count
|
||||
data = source.get("data", "")
|
||||
return (len(data) // 4, "[base64]") if data else (0, "")
|
||||
return 0, ""
|
||||
|
||||
return 0, ""
|
||||
|
||||
|
||||
_token_counter = None
|
||||
|
||||
|
||||
def get_token_counter():
|
||||
"""Get or initialize the global token counter instance.
|
||||
|
||||
Returns:
|
||||
TokenCounterBase: The token counter instance for Qwen models.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If token counter initialization fails.
|
||||
"""
|
||||
global _token_counter
|
||||
if _token_counter is None:
|
||||
# Use Qwen tokenizer for DashScope models
|
||||
# Qwen3 series uses the same tokenizer as Qwen2.5
|
||||
|
||||
# Try local tokenizer first, fall back to online if not found
|
||||
local_tokenizer_path = Path(__file__).parent.parent.parent / "tokenizer"
|
||||
|
||||
if local_tokenizer_path.exists() and (local_tokenizer_path / "tokenizer.json").exists():
|
||||
tokenizer_path = str(local_tokenizer_path)
|
||||
logger.info(f"Using local Qwen tokenizer from {tokenizer_path}")
|
||||
else:
|
||||
tokenizer_path = "Qwen/Qwen2.5-7B-Instruct"
|
||||
logger.info(
|
||||
"Local tokenizer not found, downloading from HuggingFace",
|
||||
)
|
||||
|
||||
_token_counter = HuggingFaceTokenCounter(
|
||||
pretrained_model_name_or_path=tokenizer_path,
|
||||
use_mirror=True, # Use HF mirror for users in China
|
||||
use_fast=True,
|
||||
trust_remote_code=True,
|
||||
)
|
||||
logger.debug("Token counter initialized with Qwen tokenizer")
|
||||
return _token_counter
|
||||
|
|
@ -1,17 +1,14 @@
|
|||
"""memory tools"""
|
||||
|
||||
from .base_memory_tool import BaseMemoryTool
|
||||
|
||||
# chunk tools
|
||||
from .chunk.memory_get import MemoryGet
|
||||
from .chunk.memory_search import MemorySearch
|
||||
from .delegate_task import DelegateTask
|
||||
|
||||
# history tools
|
||||
from .history.add_history import AddHistory
|
||||
from .history.read_history import ReadHistory
|
||||
from .history.read_history_v2 import ReadHistoryV2
|
||||
|
||||
# profiles tools
|
||||
from .profiles.add_draft_and_read_all_profiles import AddDraftAndReadAllProfiles
|
||||
from .profiles.add_profile import AddProfile
|
||||
|
|
@ -19,7 +16,6 @@ from .profiles.delete_profile import DeleteProfile
|
|||
from .profiles.read_all_profiles import ReadAllProfiles
|
||||
from .profiles.update_profile import UpdateProfile
|
||||
from .profiles.update_profiles_v1 import UpdateProfilesV1
|
||||
|
||||
# record tools
|
||||
from .record.add_and_retrieve_similar_memory import AddAndRetrieveSimilarMemory
|
||||
from .record.add_draft_and_retrieve_similar_memory import AddDraftAndRetrieveSimilarMemory
|
||||
|
|
|
|||
0
reme/memory/tools/file/__init__.py
Normal file
0
reme/memory/tools/file/__init__.py
Normal file
|
|
@ -16,130 +16,56 @@ Key Features:
|
|||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
from pathlib import Path
|
||||
|
||||
from agentscope.formatter import FormatterBase
|
||||
from agentscope.message import Msg, TextBlock
|
||||
from agentscope.model import ChatModelBase, OpenAIChatModel
|
||||
from agentscope.model import ChatModelBase
|
||||
from agentscope.token import HuggingFaceTokenCounter
|
||||
from agentscope.tool import Toolkit, ToolResponse
|
||||
|
||||
from .config import ReMeConfigParser
|
||||
from .core import Application
|
||||
from .memory.file_based import Compactor, Summarizer, ToolResultCompactor, ReMeInMemoryMemory, ReMeOpenAIChatFormatter, FileIO
|
||||
from .core.utils import get_hf_token_counter
|
||||
from .memory.file_based import Compactor, Summarizer, ToolResultCompactor, ReMeInMemoryMemory, ReMeOpenAIChatFormatter, \
|
||||
FileIO
|
||||
from .memory.file_based.utils import get_token_counter
|
||||
from .memory.tools import MemorySearch
|
||||
from .core.utils import load_env
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ReMeLight(Application):
|
||||
"""
|
||||
ReMe Light Application Class
|
||||
|
||||
A specialized application class that extends ReMe's core Application framework
|
||||
with advanced memory management capabilities. This class is designed to handle
|
||||
long-running conversations by providing intelligent memory compaction,
|
||||
summarization, and semantic search features.
|
||||
|
||||
Attributes:
|
||||
working_path (Path): Absolute path to the working directory for storing data
|
||||
memory_path (Path): Path to the memory storage directory
|
||||
tool_result_path (Path): Path to store large tool result files
|
||||
chat_model (ChatModelBase): Language model for generating summaries and processing
|
||||
formatter (FormatterBase): Formatter for structuring model inputs/outputs
|
||||
token_counter (HuggingFaceTokenCounter): Token counting utility for length management
|
||||
toolkit (Toolkit): Collection of tools available to the application
|
||||
max_input_length (int): Maximum allowed input length in tokens
|
||||
memory_compact_threshold (int): Threshold at which memory compaction triggers
|
||||
language (str): Language code for localization ("zh" for Chinese, empty for English)
|
||||
vector_weight (float): Weight for vector search in hybrid search (0.0-1.0)
|
||||
candidate_multiplier (float): Multiplier for candidate retrieval in search
|
||||
tool_result_threshold (int): Size threshold for tool result compaction
|
||||
retention_days (int): Number of days to retain tool result files
|
||||
summary_tasks (list[asyncio.Task]): List of background summarization tasks
|
||||
"""
|
||||
"""ReMe Light Application Class"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
working_dir: str = ".reme",
|
||||
llm_api_key: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
chat_model: ChatModelBase | None = None,
|
||||
formatter: FormatterBase | None = None,
|
||||
token_counter: HuggingFaceTokenCounter | None = None,
|
||||
toolkit: Toolkit | None = None,
|
||||
max_input_length: int = 128000,
|
||||
memory_compact_ratio: float = 0.7,
|
||||
language: str = "zh",
|
||||
vector_weight: float = 0.7,
|
||||
candidate_multiplier: float = 3.0,
|
||||
tool_result_threshold: int = 1000,
|
||||
retention_days: int = 7,
|
||||
self,
|
||||
working_dir: str = ".reme",
|
||||
llm_api_key: str | None = None,
|
||||
llm_base_url: str | None = None,
|
||||
embedding_api_key: str | None = None,
|
||||
embedding_base_url: str | None = None,
|
||||
default_as_llm_config: dict | None = None,
|
||||
default_embedding_model_config: dict | None = None,
|
||||
default_file_store_config: dict | None = None,
|
||||
vector_weight: float = 0.7,
|
||||
candidate_multiplier: float = 3.0,
|
||||
tool_result_threshold: int = 1000,
|
||||
retention_days: int = 7,
|
||||
):
|
||||
# Initialize working directory structure
|
||||
# All application data will be stored under this path
|
||||
self.working_path = Path(working_dir).absolute()
|
||||
self.working_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Create memory storage directory for persistent memory files
|
||||
self.memory_path = self.working_path / "memory"
|
||||
self.memory_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Create tool result directory for storing large tool outputs
|
||||
self.tool_result_path = self.working_path / "tool_result"
|
||||
self.tool_result_path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Apply initial parameter configuration
|
||||
self.update_params(
|
||||
max_input_length=max_input_length,
|
||||
memory_compact_ratio=memory_compact_ratio,
|
||||
language=language,
|
||||
)
|
||||
|
||||
# Store configuration parameters
|
||||
self.vector_weight: float = vector_weight
|
||||
self.candidate_multiplier: float = candidate_multiplier
|
||||
self.tool_result_threshold: int = tool_result_threshold
|
||||
self.retention_days: int = retention_days
|
||||
|
||||
load_env()
|
||||
|
||||
llm_model_name = self._safe_str("LLM_MODEL_NAME", "")
|
||||
embedding_model_name = self._safe_str("EMBEDDING_MODEL_NAME", "")
|
||||
embedding_dimensions = self._safe_int("EMBEDDING_DIMENSIONS", 1024)
|
||||
embedding_cache_enabled = self._safe_str("EMBEDDING_CACHE_ENABLED", "true").lower() == "true"
|
||||
embedding_max_cache_size = self._safe_int("EMBEDDING_MAX_CACHE_SIZE", 2000)
|
||||
embedding_max_input_length = self._safe_int("EMBEDDING_MAX_INPUT_LENGTH", 8192)
|
||||
embedding_max_batch_size = self._safe_int("EMBEDDING_MAX_BATCH_SIZE", 10)
|
||||
|
||||
# Determine if vector search should be enabled based on configuration
|
||||
# Vector search requires either an API key or a local model name
|
||||
vector_enabled = bool(embedding_api_key) or bool(embedding_model_name)
|
||||
if vector_enabled:
|
||||
logger.info("Vector search enabled.")
|
||||
else:
|
||||
logger.warning(
|
||||
"Vector search disabled. Memory search functionality will be restricted. "
|
||||
"To enable, configure: EMBEDDING_API_KEY, EMBEDDING_BASE_URL, EMBEDDING_MODEL_NAME.",
|
||||
)
|
||||
|
||||
# Check if full-text search (FTS) is enabled via environment variable
|
||||
fts_enabled = os.environ.get("FTS_ENABLED", "true").lower() == "true"
|
||||
|
||||
# Determine the memory store backend to use
|
||||
# "auto" selects based on platform (local for Windows, chroma otherwise)
|
||||
memory_store_backend = os.environ.get("MEMORY_STORE_BACKEND", "auto")
|
||||
if memory_store_backend == "auto":
|
||||
memory_backend = "local" if platform.system() == "Windows" else "chroma"
|
||||
else:
|
||||
memory_backend = memory_store_backend
|
||||
|
||||
# Initialize the parent Application class with comprehensive configuration
|
||||
super().__init__(
|
||||
llm_api_key=llm_api_key,
|
||||
|
|
@ -151,21 +77,9 @@ class ReMeLight(Application):
|
|||
enable_logo=False,
|
||||
log_to_console=False,
|
||||
parser=ReMeConfigParser,
|
||||
default_embedding_model_config={
|
||||
"model_name": embedding_model_name,
|
||||
"dimensions": embedding_dimensions,
|
||||
"enable_cache": embedding_cache_enabled,
|
||||
"use_dimensions": False,
|
||||
"max_cache_size": embedding_max_cache_size,
|
||||
"max_input_length": embedding_max_input_length,
|
||||
"max_batch_size": embedding_max_batch_size,
|
||||
},
|
||||
default_file_store_config={
|
||||
"backend": memory_backend,
|
||||
"store_name": "copaw",
|
||||
"vector_enabled": vector_enabled,
|
||||
"fts_enabled": fts_enabled,
|
||||
},
|
||||
default_as_llm_config=default_as_llm_config,
|
||||
default_embedding_model_config=default_embedding_model_config,
|
||||
default_file_store_config=default_file_store_config,
|
||||
default_file_watcher_config={
|
||||
"watch_paths": [
|
||||
str(self.working_path / "MEMORY.md"),
|
||||
|
|
@ -175,107 +89,12 @@ class ReMeLight(Application):
|
|||
},
|
||||
)
|
||||
|
||||
if chat_model is not None:
|
||||
self.chat_model: ChatModelBase = chat_model
|
||||
else:
|
||||
# add more params later
|
||||
self.chat_model = OpenAIChatModel(
|
||||
api_key=os.environ["LLM_API_KEY"],
|
||||
client_kwargs={"base_url": os.environ["LLM_BASE_URL"]},
|
||||
model_name=llm_model_name,
|
||||
)
|
||||
|
||||
if token_counter is not None:
|
||||
self.token_counter: HuggingFaceTokenCounter = token_counter
|
||||
else:
|
||||
self.token_counter = get_token_counter()
|
||||
|
||||
if formatter is not None:
|
||||
self.formatter: FormatterBase = formatter
|
||||
else:
|
||||
self.formatter = ReMeOpenAIChatFormatter(token_counter=self.token_counter)
|
||||
self.toolkit: Toolkit | None = toolkit
|
||||
|
||||
# Initialize list to track background summarization tasks
|
||||
self.summary_tasks: list[asyncio.Task] = []
|
||||
|
||||
def update_params(
|
||||
self,
|
||||
max_input_length: int,
|
||||
memory_compact_ratio: float,
|
||||
language: str,
|
||||
):
|
||||
"""
|
||||
Update runtime parameters for memory management.
|
||||
|
||||
This method allows dynamic adjustment of memory-related parameters during
|
||||
runtime. It recalculates the memory compaction threshold based on the
|
||||
new input length and compaction ratio.
|
||||
|
||||
Args:
|
||||
max_input_length (int): New maximum input length in tokens
|
||||
memory_compact_ratio (float): Ratio at which to trigger compaction (0.0-1.0)
|
||||
language (str): Language code for localization ("zh" or other)
|
||||
|
||||
Note:
|
||||
The memory_compact_threshold is calculated as:
|
||||
max_input_length * memory_compact_ratio * 0.9
|
||||
The 0.9 factor provides a safety margin before reaching the absolute limit
|
||||
"""
|
||||
# Update the maximum allowed input length
|
||||
self.max_input_length = max_input_length
|
||||
|
||||
# Calculate compaction threshold with safety margin
|
||||
# This ensures compaction happens before hitting the hard limit
|
||||
self.memory_compact_threshold = int(max_input_length * memory_compact_ratio * 0.9)
|
||||
|
||||
# Set language for localization
|
||||
if language == "zh":
|
||||
self.language = "zh"
|
||||
else:
|
||||
self.language = ""
|
||||
|
||||
@staticmethod
|
||||
def _safe_str(key: str, default: str) -> str:
|
||||
"""
|
||||
Safely retrieve a string value from an environment variable.
|
||||
|
||||
Args:
|
||||
key (str): The name of the environment variable to retrieve
|
||||
default (str): The default value to return if the variable is not set
|
||||
|
||||
Returns:
|
||||
str: The value of the environment variable, or the default if not set
|
||||
"""
|
||||
return os.environ.get(key, default)
|
||||
|
||||
@staticmethod
|
||||
def _safe_int(key: str, default: int) -> int:
|
||||
"""
|
||||
Safely retrieve an integer value from an environment variable.
|
||||
|
||||
This method handles cases where the environment variable is not set
|
||||
or contains a non-integer value by returning the specified default.
|
||||
|
||||
Args:
|
||||
key (str): The name of the environment variable to retrieve
|
||||
default (int): The default value to return on failure or if not set
|
||||
|
||||
Returns:
|
||||
int: The integer value of the environment variable, or the default
|
||||
|
||||
Note:
|
||||
Logs a warning if the value exists but cannot be parsed as an integer
|
||||
"""
|
||||
value = os.environ.get(key)
|
||||
if value is None:
|
||||
return default
|
||||
|
||||
try:
|
||||
return int(value)
|
||||
except ValueError:
|
||||
logger.warning(f"Invalid int value '{value}' for key '{key}', using default {default}")
|
||||
return default
|
||||
def calculate_memory_compact_threshold(max_input_length: float, compact_ratio: float) -> int:
|
||||
return int(max_input_length * compact_ratio * 0.9)
|
||||
|
||||
def _cleanup_tool_results(self) -> int:
|
||||
"""
|
||||
|
|
@ -287,10 +106,6 @@ class ReMeLight(Application):
|
|||
|
||||
Returns:
|
||||
int: The number of files that were successfully deleted
|
||||
|
||||
Note:
|
||||
Exceptions during cleanup are logged but do not raise errors,
|
||||
ensuring the application continues to function even if cleanup fails
|
||||
"""
|
||||
try:
|
||||
# Create a compactor instance with current configuration
|
||||
|
|
@ -307,67 +122,18 @@ class ReMeLight(Application):
|
|||
return 0
|
||||
|
||||
async def start(self):
|
||||
"""
|
||||
Start the application lifecycle.
|
||||
|
||||
This method initializes the application by calling the parent class's
|
||||
start method and performs initial cleanup of expired tool result files.
|
||||
|
||||
Returns:
|
||||
The result from the parent class's start method
|
||||
|
||||
Note:
|
||||
Tool result cleanup runs after successful startup to ensure
|
||||
the application is fully initialized before performing maintenance
|
||||
"""
|
||||
# Initialize parent application components
|
||||
"""Start the application lifecycle."""
|
||||
result = await super().start()
|
||||
# Perform initial cleanup of old tool result files
|
||||
self._cleanup_tool_results()
|
||||
return result
|
||||
|
||||
async def close(self) -> bool:
|
||||
"""
|
||||
Close the application and perform cleanup.
|
||||
|
||||
This method performs final cleanup of expired tool result files before
|
||||
shutting down the application through the parent class's close method.
|
||||
|
||||
Returns:
|
||||
bool: True if shutdown was successful, False otherwise
|
||||
|
||||
Note:
|
||||
Cleanup is performed before calling parent close to ensure
|
||||
all resources are available during the cleanup process
|
||||
"""
|
||||
# Clean up tool results before shutting down
|
||||
"""Close the application and perform cleanup."""
|
||||
self._cleanup_tool_results()
|
||||
# Shutdown parent application components
|
||||
return await super().close()
|
||||
|
||||
async def compact_tool_result(
|
||||
self,
|
||||
messages: list[Msg],
|
||||
) -> list[Msg]:
|
||||
"""
|
||||
Compact tool results by truncating large outputs and saving full content to files.
|
||||
|
||||
This method processes a list of messages and identifies tool results that exceed
|
||||
the configured size threshold. Large tool outputs are truncated in the message
|
||||
list while their full content is saved to files for later retrieval.
|
||||
|
||||
Args:
|
||||
messages (list[Msg]): List of messages to process for tool result compaction
|
||||
|
||||
Returns:
|
||||
list[Msg]: The processed message list with large tool results compacted
|
||||
|
||||
Note:
|
||||
- Tool results below the threshold remain unchanged in the messages
|
||||
- Large results are replaced with truncated versions and file references
|
||||
- Expired files are cleaned up as part of the compaction process
|
||||
- If compaction fails, the original messages are returned unchanged
|
||||
"""
|
||||
async def compact_tool_result(self, messages: list[Msg]) -> list[Msg]:
|
||||
"""Compact tool results by truncating large outputs and saving full content to files."""
|
||||
try:
|
||||
# Create compactor with instance configuration
|
||||
compactor = ToolResultCompactor(
|
||||
|
|
@ -389,38 +155,30 @@ class ReMeLight(Application):
|
|||
logger.exception(f"Error compacting tool results: {e}")
|
||||
return messages
|
||||
|
||||
async def compact_memory(self, messages: list[Msg], previous_summary: str = "") -> str:
|
||||
"""
|
||||
Compact a list of messages into a condensed summary.
|
||||
|
||||
This method uses the Compactor to reduce the length of message history
|
||||
while preserving essential information. It's useful when conversation
|
||||
history approaches the maximum input length limit.
|
||||
|
||||
Args:
|
||||
messages (list[Msg]): The list of messages to compact
|
||||
previous_summary (str): Optional previous summary to incorporate
|
||||
into the compaction process for continuity
|
||||
|
||||
Returns:
|
||||
str: A compacted summary of the messages, or empty string on failure
|
||||
|
||||
Note:
|
||||
- Compaction uses the configured language model to generate summaries
|
||||
- The compaction threshold determines when compaction is triggered
|
||||
- If compaction fails, an empty string is returned
|
||||
"""
|
||||
async def compact_memory(
|
||||
self,
|
||||
messages: list[Msg],
|
||||
as_llm: str | ChatModelBase = "default",
|
||||
as_llm_formatter: str | FormatterBase = "default",
|
||||
token_counter: HuggingFaceTokenCounter | None = None,
|
||||
language: str = "zh",
|
||||
max_input_length: float = 128 * 1024,
|
||||
compact_ratio: float = 0.7,
|
||||
previous_summary: str = "",
|
||||
) -> str:
|
||||
"""Compact a list of messages into a condensed summary."""
|
||||
try:
|
||||
# Initialize compactor with current configuration
|
||||
if token_counter is None:
|
||||
token_counter = get_hf_token_counter()
|
||||
|
||||
compactor = Compactor(
|
||||
memory_compact_threshold=self.memory_compact_threshold,
|
||||
chat_model=self.chat_model,
|
||||
formatter=self.formatter,
|
||||
token_counter=self.token_counter,
|
||||
language=self.language,
|
||||
memory_compact_threshold=self.calculate_memory_compact_threshold(max_input_length, compact_ratio),
|
||||
as_llm=as_llm,
|
||||
as_llm_formatter=as_llm_formatter,
|
||||
token_counter=token_counter,
|
||||
language=language if language == "zh" else "",
|
||||
)
|
||||
|
||||
# Execute compaction with optional previous summary context
|
||||
return await compactor.call(
|
||||
messages=messages,
|
||||
previous_summary=previous_summary,
|
||||
|
|
@ -433,25 +191,7 @@ class ReMeLight(Application):
|
|||
return ""
|
||||
|
||||
async def summary_memory(self, messages: list[Msg]) -> str:
|
||||
"""
|
||||
Generate a comprehensive summary of the given messages.
|
||||
|
||||
This method uses the Summarizer to create a detailed summary of the
|
||||
conversation history, which can be stored as persistent memory. Unlike
|
||||
compaction, summarization aims to capture key information in a format
|
||||
suitable for long-term storage and retrieval.
|
||||
|
||||
Args:
|
||||
messages (list[Msg]): The list of messages to summarize
|
||||
|
||||
Returns:
|
||||
str: A generated summary of the messages, or empty string on failure
|
||||
|
||||
Note:
|
||||
- Summarization may use tools from the toolkit to enhance the summary
|
||||
- The summary is typically stored in the memory directory
|
||||
- If summarization fails, an empty string is returned
|
||||
"""
|
||||
"""Generate a comprehensive summary of the given messages."""
|
||||
try:
|
||||
# Create toolkit if not provided
|
||||
if self.toolkit is not None:
|
||||
|
|
@ -651,24 +391,10 @@ class ReMeLight(Application):
|
|||
],
|
||||
)
|
||||
|
||||
def get_in_memory_memory(self):
|
||||
"""
|
||||
Create and return an in-memory memory instance.
|
||||
@staticmethod
|
||||
def get_in_memory_memory(token_counter: HuggingFaceTokenCounter | None = None):
|
||||
"""Create and return an in-memory memory instance."""
|
||||
if token_counter is None:
|
||||
token_counter = get_hf_token_counter()
|
||||
|
||||
This method instantiates a ReMeInMemoryMemory object configured with
|
||||
the current application's token counter, formatter, and input length limits.
|
||||
The in-memory memory provides fast, temporary storage for conversation
|
||||
context without persistence.
|
||||
|
||||
Returns:
|
||||
ReMeInMemoryMemory: A configured in-memory memory instance ready
|
||||
for storing and retrieving conversation messages
|
||||
|
||||
Note:
|
||||
- In-memory memory is volatile and cleared when the instance is destroyed
|
||||
- Useful for managing conversation context within a single session
|
||||
- Shares the same token counter as the main application
|
||||
"""
|
||||
return ReMeInMemoryMemory(
|
||||
token_counter=self.token_counter,
|
||||
)
|
||||
return ReMeInMemoryMemory(token_counter=token_counter)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue