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__)