diff --git a/poetry.lock b/poetry.lock index 0df9b1ddc..d218a9b42 100644 --- a/poetry.lock +++ b/poetry.lock @@ -2996,6 +2996,23 @@ dev = ["atheris ; python_version < \"3.12\"", "black", "mypy (==0.931)", "nox", docs = ["sphinx", "sphinx-argparse"] image = ["Pillow"] +[[package]] +name = "pika" +version = "1.3.2" +description = "Pika Python AMQP Client Library" +optional = false +python-versions = ">=3.7" +groups = ["main"] +files = [ + {file = "pika-1.3.2-py3-none-any.whl", hash = "sha256:0779a7c1fafd805672796085560d290213a465e4f6f76a6fb19e378d8041a14f"}, + {file = "pika-1.3.2.tar.gz", hash = "sha256:b2a327ddddf8570b4965b3576ac77091b850262d34ce8c1d8cb4e4146aa4145f"}, +] + +[package.extras] +gevent = ["gevent"] +tornado = ["tornado"] +twisted = ["twisted"] + [[package]] name = "pillow" version = "11.2.1" diff --git a/pyproject.toml b/pyproject.toml index fe4bc0804..f09efe38d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,6 +26,7 @@ fastapi = {extras = ["all"], version = "^0.115.12"} sentence-transformers = "^4.1.0" sqlalchemy = "^2.0.41" redis = "^6.2.0" +pika = "^1.3.2" schedule = "^1.2.2" [tool.poetry.group.dev] diff --git a/src/memos/api/config.py b/src/memos/api/config.py index 840f48613..3b62360fc 100644 --- a/src/memos/api/config.py +++ b/src/memos/api/config.py @@ -264,7 +264,7 @@ def create_user_config(user_name: str, user_id: str) -> tuple[MOSConfig, General "user": neo4j_config["user"], "password": neo4j_config["password"], "db_name": os.getenv( - "NEO4J_DB_NAME", f"db{user_id.replace('-', '')}" + "NEO4J_DB_NAME", f"memos{user_id.replace('-', '')}" ), # , replace with "auto_create": neo4j_config["auto_create"], }, diff --git a/src/memos/api/product_api.py b/src/memos/api/product_api.py index 8823ffc91..8c8f64403 100644 --- a/src/memos/api/product_api.py +++ b/src/memos/api/product_api.py @@ -26,5 +26,8 @@ if __name__ == "__main__": import uvicorn - - uvicorn.run(app, host="0.0.0.0", port=8001) + import argparse + parser = argparse.ArgumentParser() + parser.add_argument("--port", type=int, default=8001) + args = parser.parse_args() + uvicorn.run(app, host="0.0.0.0", port=args.port) diff --git a/src/memos/api/product_models.py b/src/memos/api/product_models.py index f8008c708..0e9d5ff59 100644 --- a/src/memos/api/product_models.py +++ b/src/memos/api/product_models.py @@ -150,3 +150,10 @@ class SearchRequest(BaseRequest): user_id: str = Field(..., description="User ID") query: str = Field(..., description="Search query") mem_cube_id: str | None = Field(None, description="Cube ID to search in") + + +class SuggestionRequest(BaseRequest): + """Request model for getting suggestion queries.""" + + user_id: str = Field(..., description="User ID") + language: Literal["zh", "en"] = Field("zh", description="Language for suggestions") diff --git a/src/memos/api/routers/product_router.py b/src/memos/api/routers/product_router.py index 841954947..c5a5b6e73 100644 --- a/src/memos/api/routers/product_router.py +++ b/src/memos/api/routers/product_router.py @@ -15,6 +15,7 @@ SearchRequest, SearchResponse, SimpleResponse, + SuggestionRequest, SuggestionResponse, UserRegisterRequest, UserRegisterResponse, @@ -36,6 +37,7 @@ def get_mos_product_instance(): global MOS_PRODUCT_INSTANCE if MOS_PRODUCT_INSTANCE is None: default_config = APIConfig.get_product_default_config() + print(default_config) from memos.configs.mem_os import MOSConfig mos_config = MOSConfig(**default_config) @@ -85,7 +87,6 @@ async def register_user(user_req: UserRegisterRequest): logger.error(f"Failed to register user: {traceback.format_exc()}") raise HTTPException(status_code=500, detail=str(traceback.format_exc())) from err - @router.get( "/suggestions/{user_id}", summary="Get suggestion queries", response_model=SuggestionResponse ) @@ -104,6 +105,25 @@ async def get_suggestion_queries(user_id: str): raise HTTPException(status_code=500, detail=str(traceback.format_exc())) from err +@router.post("/suggestions", summary="Get suggestion queries with language", response_model=SuggestionResponse) +async def get_suggestion_queries_post(suggestion_req: SuggestionRequest): + """Get suggestion queries for a specific user with language preference.""" + try: + mos_product = get_mos_product_instance() + suggestions = mos_product.get_suggestion_query( + user_id=suggestion_req.user_id, + language=suggestion_req.language + ) + return SuggestionResponse( + message="Suggestions retrieved successfully", data={"query": suggestions} + ) + except ValueError as err: + raise HTTPException(status_code=404, detail=str(traceback.format_exc())) from err + except Exception as err: + logger.error(f"Failed to get suggestions: {traceback.format_exc()}") + raise HTTPException(status_code=500, detail=str(traceback.format_exc())) from err + + @router.post("/get_all", summary="Get all memories for user", response_model=MemoryResponse) async def get_all_memories(memory_req: GetMemoryRequest): """Get all memories for a specific user.""" @@ -177,15 +197,19 @@ async def chat(chat_req: ChatRequest): try: mos_product = get_mos_product_instance() - def generate_chat_response(): + async def generate_chat_response(): """Generate chat response as SSE stream.""" try: - yield from mos_product.chat_with_references( + import asyncio + + for chunk in mos_product.chat_with_references( query=chat_req.query, user_id=chat_req.user_id, cube_id=chat_req.mem_cube_id, history=chat_req.history, - ) + ): + yield chunk + await asyncio.sleep(0.05) # 50ms delay between chunks except Exception as e: logger.error(f"Error in chat stream: {e}") error_data = f"data: {json.dumps({'type': 'error', 'content': str(traceback.format_exc())})}\n\n" diff --git a/src/memos/llms/hf.py b/src/memos/llms/hf.py index 895d42627..e2832f6d5 100644 --- a/src/memos/llms/hf.py +++ b/src/memos/llms/hf.py @@ -1,4 +1,5 @@ import torch +from collections.abc import Generator from transformers import ( AutoModelForCausalLM, @@ -71,6 +72,24 @@ def generate(self, messages: MessageList, past_key_values: DynamicCache | None = else: return self._generate_with_cache(prompt, past_key_values) + def generate_stream(self, messages: MessageList, past_key_values: DynamicCache | None = None) -> Generator[str, None, None]: + """ + Generate a streaming response from the model. + Args: + messages (MessageList): Chat messages for prompt construction. + past_key_values (DynamicCache | None): Optional KV cache for fast generation. + Yields: + str: Streaming model response chunks. + """ + prompt = self.tokenizer.apply_chat_template( + messages, tokenize=False, add_generation_prompt=self.config.add_generation_prompt + ) + logger.info(f"HFLLM streaming prompt: {prompt}") + if past_key_values is None: + yield from self._generate_full_stream(prompt) + else: + yield from self._generate_with_cache_stream(prompt, past_key_values) + def _generate_full(self, prompt: str) -> str: """ Generate output from scratch using the full prompt. @@ -104,6 +123,71 @@ def _generate_full(self, prompt: str) -> str: else response ) + def _generate_full_stream(self, prompt: str) -> Generator[str, None, None]: + """ + Generate output from scratch using the full prompt with streaming. + Args: + prompt (str): The input prompt string. + Yields: + str: Streaming response chunks. + """ + inputs = self.tokenizer([prompt], return_tensors="pt").to(self.model.device) + + # Get generation parameters + max_new_tokens = getattr(self.config, "max_tokens", 128) + do_sample = getattr(self.config, "do_sample", True) + remove_think_prefix = getattr(self.config, "remove_think_prefix", False) + + # Manual streaming generation + input_length = inputs.input_ids.shape[1] + generated_ids = inputs.input_ids.clone() + accumulated_text = "" + + for _ in range(max_new_tokens): + # Forward pass + with torch.no_grad(): + outputs = self.model( + input_ids=generated_ids, + use_cache=True, + return_dict=True, + ) + + # Get next token logits + next_token_logits = outputs.logits[:, -1, :] + + # Apply logits processors if sampling + if do_sample: + batch_size, _ = next_token_logits.size() + dummy_ids = torch.zeros((batch_size, 1), dtype=torch.long, device=next_token_logits.device) + filtered_logits = self.logits_processors(dummy_ids, next_token_logits) + probs = torch.softmax(filtered_logits, dim=-1) + next_token = torch.multinomial(probs, num_samples=1) + else: + next_token = torch.argmax(next_token_logits, dim=-1, keepdim=True) + + # Check for EOS token + if self._should_stop(next_token): + break + + # Append new token + generated_ids = torch.cat([generated_ids, next_token], dim=-1) + + # Decode and yield the new token + new_token_text = self.tokenizer.decode(next_token[0], skip_special_tokens=True) + if new_token_text: # Only yield non-empty tokens + accumulated_text += new_token_text + + # Apply thinking tag removal if enabled + if remove_think_prefix: + processed_text = remove_thinking_tags(accumulated_text) + # Only yield the difference (new content) + if len(processed_text) > len(accumulated_text) - len(new_token_text): + yield processed_text[len(accumulated_text) - len(new_token_text):] + else: + yield new_token_text + else: + yield new_token_text + def _generate_with_cache(self, query: str, kv: DynamicCache) -> str: """ Generate output incrementally using an existing KV cache. @@ -137,6 +221,68 @@ def _generate_with_cache(self, query: str, kv: DynamicCache) -> str: else response ) + def _generate_with_cache_stream(self, query: str, kv: DynamicCache) -> Generator[str, None, None]: + """ + Generate output incrementally using an existing KV cache with streaming. + Args: + query (str): The new user query string. + kv (DynamicCache): The prefilled KV cache. + Yields: + str: Streaming response chunks. + """ + query_ids = self.tokenizer( + query, return_tensors="pt", add_special_tokens=False + ).input_ids.to(self.model.device) + + max_new_tokens = getattr(self.config, "max_tokens", 128) + do_sample = getattr(self.config, "do_sample", True) + remove_think_prefix = getattr(self.config, "remove_think_prefix", False) + + # Initial forward pass + logits, kv = self._prefill(query_ids, kv) + next_token = self._select_next_token(logits) + + # Yield first token + first_token_text = self.tokenizer.decode(next_token[0], skip_special_tokens=True) + accumulated_text = "" + if first_token_text: + accumulated_text += first_token_text + if remove_think_prefix: + processed_text = remove_thinking_tags(accumulated_text) + if len(processed_text) > len(accumulated_text) - len(first_token_text): + yield processed_text[len(accumulated_text) - len(first_token_text):] + else: + yield first_token_text + else: + yield first_token_text + + generated = [next_token] + + # Continue generation + for _ in range(max_new_tokens - 1): + if self._should_stop(next_token): + break + logits, kv = self._prefill(next_token, kv) + next_token = self._select_next_token(logits) + + # Decode and yield the new token + new_token_text = self.tokenizer.decode(next_token[0], skip_special_tokens=True) + if new_token_text: + accumulated_text += new_token_text + + # Apply thinking tag removal if enabled + if remove_think_prefix: + processed_text = remove_thinking_tags(accumulated_text) + # Only yield the difference (new content) + if len(processed_text) > len(accumulated_text) - len(new_token_text): + yield processed_text[len(accumulated_text) - len(new_token_text):] + else: + yield new_token_text + else: + yield new_token_text + + generated.append(next_token) + @torch.no_grad() def _prefill( self, input_ids: torch.Tensor, kv: DynamicCache diff --git a/src/memos/mem_cube/general.py b/src/memos/mem_cube/general.py index d44aab915..22493428c 100644 --- a/src/memos/mem_cube/general.py +++ b/src/memos/mem_cube/general.py @@ -1,4 +1,5 @@ import os +from typing import Literal, Optional from memos.configs.mem_cube import GeneralMemCubeConfig from memos.configs.utils import get_json_file_model_schema @@ -37,10 +38,17 @@ def __init__(self, config: GeneralMemCubeConfig): else None ) - def load(self, dir: str) -> None: + def load( + self, + dir: str, + memory_types: Optional[list[Literal["text_mem", "act_mem", "para_mem"]]] = None + ) -> None: """Load memories. Args: dir (str): The directory containing the memory files. + memory_types (list[str], optional): List of memory types to load. + If None, loads all available memory types. + Options: ["text_mem", "act_mem", "para_mem"] """ loaded_schema = get_json_file_model_schema(os.path.join(dir, self.config.config_filename)) if loaded_schema != self.config.model_schema: @@ -48,35 +56,76 @@ def load(self, dir: str) -> None: f"Configuration schema mismatch. Expected {self.config.model_schema}, " f"but found {loaded_schema}." ) - self.text_mem.load(dir) if self.text_mem else None - self.act_mem.load(dir) if self.act_mem else None - self.para_mem.load(dir) if self.para_mem else None - - logger.info(f"MemCube loaded successfully from {dir}") - - def dump(self, dir: str) -> None: + + # If no specific memory types specified, load all + if memory_types is None: + memory_types = ["text_mem", "act_mem", "para_mem"] + + # Load specified memory types + if "text_mem" in memory_types and self.text_mem: + self.text_mem.load(dir) + logger.debug(f"Loaded text_mem from {dir}") + + if "act_mem" in memory_types and self.act_mem: + self.act_mem.load(dir) + logger.info(f"Loaded act_mem from {dir}") + + if "para_mem" in memory_types and self.para_mem: + self.para_mem.load(dir) + logger.info(f"Loaded para_mem from {dir}") + + logger.info(f"MemCube loaded successfully from {dir} (types: {memory_types})") + + def dump( + self, + dir: str, + memory_types: Optional[list[Literal["text_mem", "act_mem", "para_mem"]]] = None + ) -> None: """Dump memories. Args: dir (str): The directory where the memory files will be saved. + memory_types (list[str], optional): List of memory types to dump. + If None, dumps all available memory types. + Options: ["text_mem", "act_mem", "para_mem"] """ if os.path.exists(dir) and os.listdir(dir): raise MemCubeError( f"Directory {dir} is not empty. Please provide an empty directory for dumping." ) + # Always dump config self.config.to_json_file(os.path.join(dir, self.config.config_filename)) - self.text_mem.dump(dir) if self.text_mem else None - self.act_mem.dump(dir) if self.act_mem else None - self.para_mem.dump(dir) if self.para_mem else None - - logger.info(f"MemCube dumped successfully to {dir}") + + # If no specific memory types specified, dump all + if memory_types is None: + memory_types = ["text_mem", "act_mem", "para_mem"] + + # Dump specified memory types + if "text_mem" in memory_types and self.text_mem: + self.text_mem.dump(dir) + logger.info(f"Dumped text_mem to {dir}") + + if "act_mem" in memory_types and self.act_mem: + self.act_mem.dump(dir) + logger.info(f"Dumped act_mem to {dir}") + + if "para_mem" in memory_types and self.para_mem: + self.para_mem.dump(dir) + logger.info(f"Dumped para_mem to {dir}") + + logger.info(f"MemCube dumped successfully to {dir} (types: {memory_types})") @staticmethod - def init_from_dir(dir: str) -> "GeneralMemCube": + def init_from_dir( + dir: str, + memory_types: Optional[list[Literal["text_mem", "act_mem", "para_mem"]]] = None + ) -> "GeneralMemCube": """Create a MemCube instance from a MemCube directory. Args: dir (str): The directory containing the memory files. + memory_types (list[str], optional): List of memory types to load. + If None, loads all available memory types. Returns: MemCube: An instance of MemCube loaded with memories from the specified directory. @@ -84,24 +133,28 @@ def init_from_dir(dir: str) -> "GeneralMemCube": config_path = os.path.join(dir, "config.json") config = GeneralMemCubeConfig.from_json_file(config_path) mem_cube = GeneralMemCube(config) - mem_cube.load(dir) + mem_cube.load(dir, memory_types) return mem_cube @staticmethod def init_from_remote_repo( - cube_id: str, base_url: str = "https://huggingface.co/datasets" + cube_id: str, + base_url: str = "https://huggingface.co/datasets", + memory_types: Optional[list[Literal["text_mem", "act_mem", "para_mem"]]] = None ) -> "GeneralMemCube": """Create a MemCube instance from a remote repository. Args: - repo (str): The repository name. + cube_id (str): The repository name. base_url (str): The base URL of the remote repository. + memory_types (list[str], optional): List of memory types to load. + If None, loads all available memory types. Returns: MemCube: An instance of MemCube loaded with memories from the specified remote repository. """ dir = download_repo(cube_id, base_url) - return GeneralMemCube.init_from_dir(dir) + return GeneralMemCube.init_from_dir(dir, memory_types) @property def text_mem(self) -> "BaseTextMemory | None": diff --git a/src/memos/mem_os/core.py b/src/memos/mem_os/core.py index 9d8e23a0d..91ca9a236 100644 --- a/src/memos/mem_os/core.py +++ b/src/memos/mem_os/core.py @@ -1,5 +1,5 @@ import os - +import uuid from datetime import datetime from pathlib import Path from threading import Lock @@ -539,7 +539,7 @@ def add( memories = self.mem_reader.get_memory( messages_list, type="chat", - info={"user_id": target_user_id, "session_id": self.session_id}, + info={"user_id": target_user_id, "session_id": str(uuid.uuid4())}, ) for mem in memories: self.mem_cubes[mem_cube_id].text_mem.add(mem) @@ -568,7 +568,7 @@ def add( memories = self.mem_reader.get_memory( messages_list, type="chat", - info={"user_id": target_user_id, "session_id": self.session_id}, + info={"user_id": target_user_id, "session_id": str(uuid.uuid4())}, ) for mem in memories: self.mem_cubes[mem_cube_id].text_mem.add(mem) @@ -581,7 +581,7 @@ def add( doc_memory = self.mem_reader.get_memory( documents, type="doc", - info={"user_id": target_user_id, "session_id": self.session_id}, + info={"user_id": target_user_id, "session_id": str(uuid.uuid4())}, ) for mem in doc_memory: self.mem_cubes[mem_cube_id].text_mem.add(mem) diff --git a/src/memos/mem_os/product.py b/src/memos/mem_os/product.py index 4c551640a..7ec538fcf 100644 --- a/src/memos/mem_os/product.py +++ b/src/memos/mem_os/product.py @@ -126,7 +126,11 @@ def _restore_user_instances(self) -> None: # Store user config and cube access self.user_configs[user_id] = config self._load_user_cube_access(user_id) - logger.info(f"Restored user configuration for {user_id}") + + # Pre-load all cubes for this user + self._preload_user_cubes(user_id) + + logger.info(f"Restored user configuration and pre-loaded cubes for {user_id}") except Exception as e: logger.error(f"Failed to restore user configuration for {user_id}: {e}") @@ -134,6 +138,38 @@ def _restore_user_instances(self) -> None: except Exception as e: logger.error(f"Error during user instance restoration: {e}") + def _preload_user_cubes(self, user_id: str) -> None: + """Pre-load all cubes for a user into memory. + + Args: + user_id (str): The user ID to pre-load cubes for. + """ + try: + # Get user's accessible cubes from persistent storage + accessible_cubes = self.global_user_manager.get_user_cubes(user_id) + + for cube in accessible_cubes: + if cube.cube_id not in self.mem_cubes: + try: + if cube.cube_path and os.path.exists(cube.cube_path): + # Pre-load cube with all memory types + self.register_mem_cube( + cube.cube_path, + cube.cube_id, + user_id, + memory_types=["act_mem"] if self.config.enable_activation_memory else [] + ) + logger.info(f"Pre-loaded cube {cube.cube_id} for user {user_id}") + else: + logger.warning( + f"Cube path {cube.cube_path} does not exist for cube {cube.cube_id}, skipping pre-load" + ) + except Exception as e: + logger.error(f"Failed to pre-load cube {cube.cube_id} for user {user_id}: {e}") + + except Exception as e: + logger.error(f"Error pre-loading cubes for user {user_id}: {e}") + def _ensure_user_instance(self, user_id: str, max_instances: int | None = None) -> None: """ Ensure user configuration exists, creating it if necessary. @@ -245,7 +281,13 @@ def _load_user_cubes(self, user_id: str) -> None: try: if cube.cube_path and os.path.exists(cube.cube_path): # Use MOSCore's register_mem_cube method directly - self.register_mem_cube(cube.cube_path, cube.cube_id, user_id) + # Only load act_mem since text_mem is stored in database + self.register_mem_cube( + cube.cube_path, + cube.cube_id, + user_id, + memory_types=["act_mem"] + ) else: logger.warning( f"Cube path {cube.cube_path} does not exist for cube {cube.cube_id}" @@ -380,7 +422,6 @@ def _chunk_response_with_tiktoken( """ if self.tokenizer: # Use tiktoken for proper token-based chunking - print(response) tokens = self.tokenizer.encode(response) for i in range(0, len(tokens), chunk_size): @@ -420,28 +461,50 @@ def _send_message_to_scheduler( self.mem_scheduler.submit_messages(messages=[message_item]) def register_mem_cube( - self, mem_cube_name_or_path: str, mem_cube_id: str | None = None, user_id: str | None = None + self, + mem_cube_name_or_path_or_object: str | GeneralMemCube, + mem_cube_id: str | None = None, + user_id: str | None = None, + memory_types: list[Literal["text_mem", "act_mem", "para_mem"]] | None = None ) -> None: """ Register a MemCube with the MOS. Args: - mem_cube_name_or_path (str): The name or path of the MemCube to register. + mem_cube_name_or_path_or_object (str | GeneralMemCube): The name, path, or GeneralMemCube object to register. mem_cube_id (str, optional): The identifier for the MemCube. If not provided, a default ID is used. + user_id (str, optional): The user ID to register the cube for. + memory_types (list[str], optional): List of memory types to load. + If None, loads all available memory types. + Options: ["text_mem", "act_mem", "para_mem"] """ - - if mem_cube_id in self.mem_cubes: - logger.info(f"MemCube with ID {mem_cube_id} already in MOS, skip install.") + # Handle different input types + if isinstance(mem_cube_name_or_path_or_object, GeneralMemCube): + # Direct GeneralMemCube object provided + mem_cube = mem_cube_name_or_path_or_object + if mem_cube_id is None: + mem_cube_id = f"cube_{id(mem_cube)}" # Generate a unique ID else: + # String path provided + mem_cube_name_or_path = mem_cube_name_or_path_or_object + if mem_cube_id is None: + mem_cube_id = mem_cube_name_or_path + + if mem_cube_id in self.mem_cubes: + logger.info(f"MemCube with ID {mem_cube_id} already in MOS, skip install.") + return + + # Create MemCube from path if os.path.exists(mem_cube_name_or_path): - self.mem_cubes[mem_cube_id] = GeneralMemCube.init_from_dir(mem_cube_name_or_path) + mem_cube = GeneralMemCube.init_from_dir(mem_cube_name_or_path, memory_types) else: logger.warning( f"MemCube {mem_cube_name_or_path} does not exist, try to init from remote repo." ) - self.mem_cubes[mem_cube_id] = GeneralMemCube.init_from_remote_repo( - mem_cube_name_or_path - ) + mem_cube = GeneralMemCube.init_from_remote_repo(mem_cube_name_or_path, memory_types=memory_types) + + # Register the MemCube + self.mem_cubes[mem_cube_id] = mem_cube def user_register( self, @@ -482,7 +545,7 @@ def user_register( user_config = self._create_user_config(user_id, user_config) # Create a default cube for the user using MOSCore's methods - default_cube_name = f"{user_name}_default_cube" + default_cube_name = f"{user_name}_{user_id}_default_cube" mem_cube_name_or_path = f"{CUBE_PATH}/{default_cube_name}" default_cube_id = self.create_cube_for_user( cube_name=default_cube_name, owner_id=user_id, cube_path=mem_cube_name_or_path @@ -495,7 +558,12 @@ def user_register( print(e) # Register the default cube with MOS TODO overide - self.register_mem_cube(mem_cube_name_or_path, default_cube_id, user_id) + self.register_mem_cube( + mem_cube_name_or_path_or_object=default_mem_cube, + mem_cube_id=default_cube_id, + user_id=user_id, + memory_types=["act_mem"] if self.config.enable_activation_memory else [] + ) # Add interests to the default cube if provided if interests: @@ -511,37 +579,56 @@ def user_register( except Exception as e: return {"status": "error", "message": f"Failed to register user: {e!s}"} - def get_suggestion_query(self, user_id: str) -> list[str]: + def get_suggestion_query(self, user_id: str, language: str = "zh") -> list[str]: """Get suggestion query from LLM. Args: user_id (str): User ID. + language (str): Language for suggestions ("zh" or "en"). Returns: list[str]: The suggestion query list. """ - suggestion_prompt = """ - You are a helpful assistant that can help users to generate suggestion query - I will get some user recently memories, - you should generate some suggestion query , the query should be user what to query, - user recently memories is : - {memories} - please generate 3 suggestion query, - output should be a json format, the key is "query", the value is a list of suggestion query. - - example: - {{ - "query": ["query1", "query2", "query3"] - }} - """ - memories = "\n".join( - [ - m.memory - for m in super().search("my recently memories", user_id=user_id, top_k=10)[ - "text_mem" - ][0]["memories"] - ] - ) + if language == "zh": + suggestion_prompt = """ + 你是一个有用的助手,可以帮助用户生成建议查询。 + 我将获取用户最近的一些记忆, + 你应该生成一些建议查询,这些查询应该是用户想要查询的内容, + 用户最近的记忆是: + {memories} + 请生成3个建议查询用中文, + 输出应该是json格式,键是"query",值是一个建议查询列表。 + + 示例: + {{ + "query": ["查询1", "查询2", "查询3"] + }} + """ + else: # English + suggestion_prompt = """ + You are a helpful assistant that can help users to generate suggestion query. + I will get some user recently memories, + you should generate some suggestion query, the query should be user what to query, + user recently memories is: + {memories} + please generate 3 suggestion query in English, + output should be a json format, the key is "query", the value is a list of suggestion query. + + example: + {{ + "query": ["query1", "query2", "query3"] + }} + """ + text_mem_result = super().search("my recently memories", user_id=user_id, top_k=10)["text_mem"] + if text_mem_result: + memories = "\n".join( + [ + m.memory + for m in text_mem_result[0]["memories"] + ] + ) + else: + memories = "" message_list = [{"role": "system", "content": suggestion_prompt.format(memories=memories)}] response = self.chat_llm.generate(message_list) response_json = json.loads(response) @@ -622,9 +709,12 @@ def chat_with_references( self._load_user_cubes(user_id) time_start = time.time() - memories_list = super().search( - query, user_id, install_cube_ids=[cube_id] if cube_id else None - )["text_mem"][0]["memories"] + memories_list = [] + memories_result = super().search( + query, user_id, install_cube_ids=[cube_id] if cube_id else None, top_k=10 + )["text_mem"] + if memories_result: + memories_list = memories_result[0]["memories"] # Build custom system prompt with relevant memories system_prompt = self._build_system_prompt(user_id, memories_list) @@ -638,11 +728,12 @@ def chat_with_references( current_messages = [ {"role": "system", "content": system_prompt}, *chat_history.chat_history, - {"role": "user", "content": query}, + {"role": "user", "content": query + "/nothink"}, ] # Generate response with custom prompt past_key_values = None + response_stream = None if self.config.enable_activation_memory: # Handle activation memory (copy MOSCore logic) for mem_cube_id, mem_cube in self.mem_cubes.items(): @@ -656,9 +747,13 @@ def chat_with_references( else: logger.info("past_key_values is None will not apply to chat") break - response = self.chat_llm.generate(current_messages, past_key_values=past_key_values) + if self.config.chat_model.backend == "huggingface": + response_stream = self.chat_llm.generate_stream(current_messages, past_key_values=past_key_values) else: - response = self.chat_llm.generate(current_messages) + if self.config.chat_model.backend == "huggingface": + response_stream = self.chat_llm.generate_stream(current_messages) + else: + response_stream = self.chat_llm.generate(current_messages) time_end = time.time() @@ -666,10 +761,22 @@ def chat_with_references( # Initialize buffer for streaming buffer = "" + full_response = "" # Use tiktoken for proper token-based chunking - for chunk in self._chunk_response_with_tiktoken(response, chunk_size=5): + if self.config.chat_model.backend != "huggingface": + # For non-huggingface backends, we need to collect the full response first + full_response_text = "" + for chunk in response_stream: + if chunk in ["", ""]: + continue + full_response_text += chunk + response_stream = self._chunk_response_with_tiktoken(full_response_text, chunk_size=5) + for chunk in response_stream: + if chunk in ["", ""]: + continue buffer += chunk + full_response += chunk # Process buffer to ensure complete reference tags processed_chunk, remaining_buffer = self._process_streaming_references_complete(buffer) @@ -700,16 +807,26 @@ def chat_with_references( total_time = round(float(time_end - time_start), 1) yield f"data: {json.dumps({'type': 'time', 'data': {'total_time': total_time, 'speed_improvement': '23%'}})}\n\n" chat_history.chat_history.append({"role": "user", "content": query}) - chat_history.chat_history.append({"role": "assistant", "content": response}) + chat_history.chat_history.append({"role": "assistant", "content": full_response}) self._send_message_to_scheduler( user_id=user_id, mem_cube_id=cube_id, query=query, label=QUERY_LABEL ) self._send_message_to_scheduler( - user_id=user_id, mem_cube_id=cube_id, query=response, label=ANSWER_LABEL + user_id=user_id, mem_cube_id=cube_id, query=full_response, label=ANSWER_LABEL ) self.chat_history_manager[user_id] = chat_history yield f"data: {json.dumps({'type': 'end'})}\n\n" + self.add( + user_id=user_id, + messages=[ + {"role": "user", "content": query}, + {"role": "assistant", "content": full_response} + ], + mem_cube_id=cube_id + ) + if len(self.chat_history_manager[user_id].chat_history) > 30: + self.chat_history_manager[user_id].chat_history.pop(0) def get_all( self, @@ -742,7 +859,7 @@ def get_all( "LongTermMemory": 0.40, "UserMemory": 0.40, } - tree_result = convert_graph_to_tree_forworkmem( + tree_result, node_type_count = convert_graph_to_tree_forworkmem( memories, target_node_count=150, type_ratios=custom_type_ratios ) memories_filtered = filter_nodes_by_tree_ids(tree_result, memories) @@ -751,7 +868,7 @@ def get_all( tree_result["children"] = children_sort memories_filtered["tree_structure"] = tree_result reformat_memory_list.append( - {"cube_id": memory["cube_id"], "memories": [memories_filtered]} + {"cube_id": memory["cube_id"], "memories": [memories_filtered], "memory_statistics": node_type_count} ) elif memory_type == "act_mem": reformat_memory_list.append( @@ -815,7 +932,7 @@ def get_subgraph( for memory in memory_list: memories = remove_embedding_recursive(memory["memories"]) custom_type_ratios = {"WorkingMemory": 0.20, "LongTermMemory": 0.40, "UserMemory": 0.4} - tree_result = convert_graph_to_tree_forworkmem( + tree_result, node_type_count = convert_graph_to_tree_forworkmem( memories, target_node_count=150, type_ratios=custom_type_ratios ) memories_filtered = filter_nodes_by_tree_ids(tree_result, memories) @@ -824,7 +941,7 @@ def get_subgraph( tree_result["children"] = children_sort memories_filtered["tree_structure"] = tree_result reformat_memory_list.append( - {"cube_id": memory["cube_id"], "memories": [memories_filtered]} + {"cube_id": memory["cube_id"], "memories": [memories_filtered], "memory_statistics": node_type_count} ) return reformat_memory_list diff --git a/src/memos/mem_os/utils/format_utils.py b/src/memos/mem_os/utils/format_utils.py index ebf18ba22..c205565cb 100644 --- a/src/memos/mem_os/utils/format_utils.py +++ b/src/memos/mem_os/utils/format_utils.py @@ -508,6 +508,10 @@ def convert_graph_to_tree_forworkmem( for original_edge in original_edges: if original_edge["type"] == "PARENT": filter_original_edges.append(original_edge) + node_type_count = {} + for node in original_nodes: + node_type = node.get("metadata", {}).get("memory_type", "Unknown") + node_type_count[node_type] = node_type_count.get(node_type, 0) + 1 original_edges = filter_original_edges # Use enhanced type-balanced sampling if len(original_nodes) > target_node_count: @@ -524,11 +528,15 @@ def convert_graph_to_tree_forworkmem( node_map = {} for node in nodes: memory = node.get("memory", "") + node_name = extract_node_name(memory) + memory_key = node.get("metadata", {}).get("key", node_name) + usage = node.get("metadata", {}).get("usage",[]) + frequency = len(usage) node_map[node["id"]] = { "id": node["id"], "value": memory, - "frequency": random.randint(1, 100), - "node_name": extract_node_name(memory), + "frequency": frequency, + "node_name": memory_key, "memory_type": node.get("metadata", {}).get("memory_type", "Unknown"), "children": [], } @@ -621,7 +629,7 @@ def build_tree(node_id: str) -> dict[str, Any]: "frequency": 0, } - return result + return result, node_type_count def print_tree_structure(node: dict[str, Any], level: int = 0, max_level: int = 5): diff --git a/src/memos/mem_scheduler/modules/schemas.py b/src/memos/mem_scheduler/modules/schemas.py index 0927e9226..06fb813fd 100644 --- a/src/memos/mem_scheduler/modules/schemas.py +++ b/src/memos/mem_scheduler/modules/schemas.py @@ -4,7 +4,6 @@ from typing import ClassVar, TypeVar, List, Optional from uuid import uuid4 -from databricks.sdk.service.cleanrooms import ListCleanRoomsResponse from pydantic import BaseModel, Field, computed_field from typing_extensions import TypedDict diff --git a/src/memos/mem_user/persistent_user_manager.py b/src/memos/mem_user/persistent_user_manager.py index e3c476262..ff10212ea 100644 --- a/src/memos/mem_user/persistent_user_manager.py +++ b/src/memos/mem_user/persistent_user_manager.py @@ -13,7 +13,7 @@ from memos.configs.mem_os import MOSConfig from memos.log import get_logger -from memos.mem_user.user_manager import Base, UserManager +from memos.mem_user.user_manager import Base, UserManager, UserRole logger = get_logger(__name__)