mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-06 02:48:22 +00:00
fin(emb_dim) (#150)
* fix(logger): handle file logging configuration errors gracefully * fix(embedding): handle embedding dimension mismatches and improve caching * fix(file-watcher): clear file store on changes to prevent stale data * refactor(logger): update logger implementation and fix message translation * refactor(logger): update logger configuration and add documentation
This commit is contained in:
parent
083ed6a137
commit
8b1698451e
9 changed files with 192 additions and 62 deletions
|
|
@ -6,7 +6,7 @@ from . import extension
|
|||
from . import memory
|
||||
from .reme import ReMe
|
||||
|
||||
__version__ = "0.3.0.6b2"
|
||||
__version__ = "0.3.0.6b3"
|
||||
|
||||
__all__ = [
|
||||
"config",
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ class BaseEmbeddingModel(ABC):
|
|||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
model_name: str = "",
|
||||
dimensions: int | None = 1024,
|
||||
dimensions: int = 1024,
|
||||
use_dimensions: bool = False,
|
||||
max_batch_size: int = 10,
|
||||
max_retries: int = 3,
|
||||
|
|
@ -99,7 +99,34 @@ class BaseEmbeddingModel(ABC):
|
|||
"""Truncate a list of texts to max_input_length."""
|
||||
return [self._truncate_text(text) for text in texts]
|
||||
|
||||
def _get_cache_key(self, text: str) -> str:
|
||||
def _validate_and_adjust_embedding(self, embedding: list[float]) -> list[float]:
|
||||
"""Validate and adjust embedding dimensions to match expected dimensions.
|
||||
|
||||
Args:
|
||||
embedding: The embedding vector to validate
|
||||
|
||||
Returns:
|
||||
Embedding vector adjusted to match self.dimensions
|
||||
"""
|
||||
actual_len = len(embedding)
|
||||
if actual_len == self.dimensions:
|
||||
return embedding
|
||||
|
||||
elif actual_len < self.dimensions:
|
||||
logger.warning(
|
||||
f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is less than expected {self.dimensions}, "
|
||||
f"padding with zeros",
|
||||
)
|
||||
return embedding + [0.0] * (self.dimensions - actual_len)
|
||||
|
||||
else:
|
||||
logger.warning(
|
||||
f"[ACTUAL_EMB_LENGTH]Embedding dimensions {actual_len} is greater than expected {self.dimensions}, "
|
||||
f"truncating to {self.dimensions}",
|
||||
)
|
||||
return embedding[: self.dimensions]
|
||||
|
||||
def _get_cache_key(self, text: str, dimensions: int) -> str:
|
||||
"""Generate a cache key by hashing text + model_name + dimensions.
|
||||
|
||||
This ensures that the same text produces different cache keys when
|
||||
|
|
@ -107,12 +134,13 @@ class BaseEmbeddingModel(ABC):
|
|||
|
||||
Args:
|
||||
text: Input text to hash
|
||||
dimensions: Vector dimensions of the embeddings
|
||||
|
||||
Returns:
|
||||
SHA256 hash combining text, model name, and dimensions
|
||||
"""
|
||||
# Combine text, model_name, and dimensions to create unique cache key
|
||||
cache_string = f"{text}|{self.model_name}|{self.dimensions}"
|
||||
cache_string = f"{text}|{self.model_name}|{dimensions}"
|
||||
return hashlib.sha256(cache_string.encode("utf-8")).hexdigest()
|
||||
|
||||
def _get_cache_file_path(self) -> Path:
|
||||
|
|
@ -164,6 +192,13 @@ class BaseEmbeddingModel(ABC):
|
|||
if cache_key in self._embedding_cache:
|
||||
continue
|
||||
|
||||
if len(embedding) != self.dimensions:
|
||||
logger.warning(
|
||||
f"Embedding dimensions mismatch for cache key {cache_key}, "
|
||||
f"expected {self.dimensions}, got {len(embedding)}",
|
||||
)
|
||||
continue
|
||||
|
||||
# Respect max_cache_size during loading
|
||||
if len(self._embedding_cache) >= self.max_cache_size:
|
||||
logger.info(
|
||||
|
|
@ -204,6 +239,12 @@ class BaseEmbeddingModel(ABC):
|
|||
try:
|
||||
with open(cache_file, "w", encoding="utf-8") as f:
|
||||
for cache_key, embedding in self._embedding_cache.items():
|
||||
if len(embedding) != self.dimensions:
|
||||
logger.warning(
|
||||
f"Embedding dimensions mismatch for cache key {cache_key}, "
|
||||
f"expected {self.dimensions}, got {len(embedding)}",
|
||||
)
|
||||
continue
|
||||
cache_entry = {cache_key: embedding}
|
||||
f.write(json.dumps(cache_entry, ensure_ascii=False) + "\n")
|
||||
|
||||
|
|
@ -223,16 +264,27 @@ class BaseEmbeddingModel(ABC):
|
|||
if not self.enable_cache:
|
||||
return None
|
||||
|
||||
cache_key = self._get_cache_key(text)
|
||||
cache_key = self._get_cache_key(text, self.dimensions)
|
||||
if cache_key in self._embedding_cache:
|
||||
embeddings: list[float] = self._embedding_cache[cache_key]
|
||||
|
||||
# Validate embedding dimensions match expected dimensions
|
||||
if len(embeddings) != self.dimensions:
|
||||
logger.warning(
|
||||
f"Cached embedding dimensions mismatch: expected {self.dimensions}, "
|
||||
f"got {len(embeddings)}. Removing invalid cache entry.",
|
||||
)
|
||||
del self._embedding_cache[cache_key]
|
||||
self._cache_misses += 1
|
||||
return None
|
||||
|
||||
# Move to end (most recently used)
|
||||
self._embedding_cache.move_to_end(cache_key)
|
||||
self._cache_hits += 1
|
||||
text_preview = text[:50] + "..." if len(text) > 50 else text
|
||||
logger.info(
|
||||
f"Cache hit for text: '{text_preview}' (hits: {self._cache_hits}, misses: {self._cache_misses})",
|
||||
)
|
||||
return self._embedding_cache[cache_key]
|
||||
logger.info(f"Cache hit for text: {text_preview} (hits: {self._cache_hits}, misses: {self._cache_misses})")
|
||||
return embeddings
|
||||
|
||||
self._cache_misses += 1
|
||||
return None
|
||||
|
||||
|
|
@ -249,9 +301,15 @@ class BaseEmbeddingModel(ABC):
|
|||
if self.max_cache_size <= 0:
|
||||
return
|
||||
|
||||
cache_key = self._get_cache_key(text)
|
||||
cache_key = self._get_cache_key(text, self.dimensions)
|
||||
if len(embedding) != self.dimensions:
|
||||
logger.warning(
|
||||
f"[PUT_TO_CACHE] Embedding dimensions mismatch for cache key {cache_key}, "
|
||||
f"expected {self.dimensions}, got real length {len(embedding)}",
|
||||
)
|
||||
return
|
||||
|
||||
# Remove oldest entry if cache is full
|
||||
# Remove the oldest entry if cache is full
|
||||
if len(self._embedding_cache) >= self.max_cache_size and cache_key not in self._embedding_cache:
|
||||
self._embedding_cache.popitem(last=False)
|
||||
|
||||
|
|
@ -299,7 +357,7 @@ class BaseEmbeddingModel(ABC):
|
|||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = await self._get_embeddings([truncated_text], **kwargs)
|
||||
embedding = result[0]
|
||||
embedding = self._validate_and_adjust_embedding(result[0])
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
|
|
@ -345,8 +403,9 @@ class BaseEmbeddingModel(ABC):
|
|||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
results[orig_idx] = embedding
|
||||
self._put_to_cache(text, embedding)
|
||||
adjusted_embedding = self._validate_and_adjust_embedding(embedding)
|
||||
results[orig_idx] = adjusted_embedding
|
||||
self._put_to_cache(text, adjusted_embedding)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Model {self.model_name} batch failed: {e}")
|
||||
|
|
@ -371,7 +430,7 @@ class BaseEmbeddingModel(ABC):
|
|||
for i in range(self.max_retries):
|
||||
try:
|
||||
result = self._get_embeddings_sync([truncated_text], **kwargs)
|
||||
embedding = result[0]
|
||||
embedding = self._validate_and_adjust_embedding(result[0])
|
||||
# Store in cache
|
||||
self._put_to_cache(truncated_text, embedding)
|
||||
return embedding
|
||||
|
|
@ -417,8 +476,9 @@ class BaseEmbeddingModel(ABC):
|
|||
if batch_embeddings:
|
||||
# Store results and cache them
|
||||
for orig_idx, text, embedding in zip(batch_indices, batch_texts, batch_embeddings):
|
||||
results[orig_idx] = embedding
|
||||
self._put_to_cache(text, embedding)
|
||||
adjusted_embedding = self._validate_and_adjust_embedding(embedding)
|
||||
results[orig_idx] = adjusted_embedding
|
||||
self._put_to_cache(text, adjusted_embedding)
|
||||
break
|
||||
except Exception as exc:
|
||||
logger.error(f"Model {self.model_name} batch failed: {exc}")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""ChromaDB storage backend for file store."""
|
||||
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
|
|
@ -355,12 +356,41 @@ class ChromaFileStore(BaseFileStore):
|
|||
where_filter = {"source": {"$in": [s.value for s in sources]}}
|
||||
|
||||
# Perform vector search
|
||||
results = self.chunks_collection.query(
|
||||
query_embeddings=[query_embedding],
|
||||
n_results=limit,
|
||||
where=where_filter,
|
||||
include=["documents", "metadatas", "distances"],
|
||||
)
|
||||
try:
|
||||
results = self.chunks_collection.query(
|
||||
query_embeddings=[query_embedding],
|
||||
n_results=limit,
|
||||
where=where_filter,
|
||||
include=["documents", "metadatas", "distances"],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Vector search failed: {e}, falling back to random results")
|
||||
# Fallback: get some documents without vector search and assign random scores
|
||||
try:
|
||||
fallback_results = self.chunks_collection.get(
|
||||
where=where_filter,
|
||||
limit=limit,
|
||||
include=["documents", "metadatas"],
|
||||
)
|
||||
search_results = []
|
||||
if fallback_results["ids"]:
|
||||
for i, _ in enumerate(fallback_results["ids"]):
|
||||
metadata = fallback_results["metadatas"][i]
|
||||
search_results.append(
|
||||
MemorySearchResult(
|
||||
path=metadata["path"],
|
||||
start_line=metadata["start_line"],
|
||||
end_line=metadata["end_line"],
|
||||
score=random.uniform(0.3, 0.7), # Random score in middle range
|
||||
snippet=fallback_results["documents"][i],
|
||||
source=MemorySource(metadata["source"]),
|
||||
raw_metric=None,
|
||||
),
|
||||
)
|
||||
return search_results
|
||||
except Exception as fallback_e:
|
||||
logger.error(f"Fallback search also failed: {fallback_e}")
|
||||
return []
|
||||
|
||||
search_results = []
|
||||
if results["ids"] and results["ids"][0]:
|
||||
|
|
@ -430,7 +460,7 @@ class ChromaFileStore(BaseFileStore):
|
|||
# ChromaDB where_document uses $contains for substring matching (case-sensitive)
|
||||
# Use multiple case variants to improve recall
|
||||
if len(word_variants_list) == 1:
|
||||
where_document = {"$contains": word_variants_list[0]}
|
||||
where_document: dict = {"$contains": word_variants_list[0]}
|
||||
else:
|
||||
where_document = {"$or": [{"$contains": w} for w in word_variants_list]}
|
||||
|
||||
|
|
|
|||
|
|
@ -259,6 +259,8 @@ class LocalFileStore(BaseFileStore):
|
|||
if not query_embedding:
|
||||
return []
|
||||
|
||||
expected_dim = self.embedding_dim
|
||||
|
||||
# Collect candidate chunks with embeddings
|
||||
candidates = [
|
||||
chunk for chunk in self._chunks.values() if (not sources or chunk.source in sources) and chunk.embedding
|
||||
|
|
@ -267,9 +269,29 @@ class LocalFileStore(BaseFileStore):
|
|||
if not candidates:
|
||||
return []
|
||||
|
||||
# Validate and fix chunk embedding dimensions
|
||||
valid_embeddings = []
|
||||
for chunk in candidates:
|
||||
emb = chunk.embedding
|
||||
emb_len = len(emb)
|
||||
if emb_len != expected_dim:
|
||||
if emb_len < expected_dim:
|
||||
emb = emb + [0.0] * (expected_dim - emb_len)
|
||||
logger.warning(
|
||||
f"Chunk embedding dimension {emb_len} < expected {expected_dim}, "
|
||||
f"padded with zeros (chunk_id={chunk.id})",
|
||||
)
|
||||
else:
|
||||
emb = emb[:expected_dim]
|
||||
logger.warning(
|
||||
f"Chunk embedding dimension {emb_len} > expected {expected_dim}, "
|
||||
f"truncated to {expected_dim} (chunk_id={chunk.id})",
|
||||
)
|
||||
valid_embeddings.append(emb)
|
||||
|
||||
# Build embedding matrix and compute similarities in batch
|
||||
query_array = np.array([query_embedding]) # Shape: (1, emb_size)
|
||||
chunk_embeddings = np.array([chunk.embedding for chunk in candidates]) # Shape: (n, emb_size)
|
||||
chunk_embeddings = np.array(valid_embeddings) # Shape: (n, emb_size)
|
||||
similarities = batch_cosine_similarity(query_array, chunk_embeddings)[0] # Shape: (n,)
|
||||
|
||||
# Build results
|
||||
|
|
|
|||
|
|
@ -141,6 +141,7 @@ class DeltaFileWatcher(BaseFileWatcher):
|
|||
async def _on_changes(self, changes: set[tuple[Change, str]]):
|
||||
"""Handle file changes with incremental synchronization."""
|
||||
self.dirty = True
|
||||
await self.file_store.clear_all()
|
||||
|
||||
for change_type, path in changes:
|
||||
if change_type == Change.added:
|
||||
|
|
|
|||
|
|
@ -44,6 +44,8 @@ class FullFileWatcher(BaseFileWatcher):
|
|||
async def _on_changes(self, changes: set[tuple[Change, str]]):
|
||||
"""Handle file changes with full synchronization"""
|
||||
self.dirty = True
|
||||
await self.file_store.clear_all()
|
||||
|
||||
for change_type, path in changes:
|
||||
if change_type in [Change.added, Change.modified]:
|
||||
file_meta = await self._build_file_metadata(path)
|
||||
|
|
|
|||
|
|
@ -18,26 +18,6 @@ def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool
|
|||
# Remove default handler to avoid duplicate logs
|
||||
logger.remove()
|
||||
|
||||
# Ensure the logging directory exists
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
|
||||
# Generate filename based on the current timestamp
|
||||
# Use dashes instead of colons for Windows compatibility
|
||||
current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
log_filename = f"{current_ts}.log"
|
||||
log_filepath = os.path.join(log_dir, log_filename)
|
||||
|
||||
# Configure file-based logging with rotation and compression
|
||||
logger.add(
|
||||
log_filepath,
|
||||
level=level,
|
||||
rotation="00:00",
|
||||
retention="7 days",
|
||||
compression="zip",
|
||||
encoding="utf-8",
|
||||
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}",
|
||||
)
|
||||
|
||||
# Configure colorized standard output logging if enabled
|
||||
if log_to_console:
|
||||
logger.add(
|
||||
|
|
@ -46,3 +26,27 @@ def init_logger(log_dir: str = "logs", level: str = "INFO", log_to_console: bool
|
|||
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}",
|
||||
colorize=True,
|
||||
)
|
||||
|
||||
# Try to configure file-based logging (skip if permission denied)
|
||||
try:
|
||||
# Ensure the logging directory exists
|
||||
os.makedirs(log_dir, exist_ok=True)
|
||||
|
||||
# Generate filename based on the current timestamp
|
||||
# Use dashes instead of colons for Windows compatibility
|
||||
current_ts = datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
log_filename = f"{current_ts}.log"
|
||||
log_filepath = os.path.join(log_dir, log_filename)
|
||||
|
||||
# Configure file-based logging with rotation and compression
|
||||
logger.add(
|
||||
log_filepath,
|
||||
level=level,
|
||||
rotation="00:00",
|
||||
retention="7 days",
|
||||
compression="zip",
|
||||
encoding="utf-8",
|
||||
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {file}:{line} | {function} | {message}",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Error configuring file logging: {e}")
|
||||
|
|
|
|||
|
|
@ -38,7 +38,7 @@ class CustomFormatter(logging.Formatter):
|
|||
return super().format(record)
|
||||
|
||||
|
||||
def get_logger(
|
||||
def get_loggerv2(
|
||||
name: str = "reme",
|
||||
log_dir: str = "logs",
|
||||
level: str = "INFO",
|
||||
|
|
@ -80,22 +80,26 @@ def get_logger(
|
|||
|
||||
# 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)
|
||||
try:
|
||||
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)
|
||||
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)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error configuring file logging: {e}")
|
||||
|
||||
# Configure console logging
|
||||
if log_to_console:
|
||||
|
|
@ -107,3 +111,10 @@ def get_logger(
|
|||
# Cache logger
|
||||
_loggers[name] = logger
|
||||
return logger
|
||||
|
||||
|
||||
def get_logger():
|
||||
"""Get a configured logger instance using loguru."""
|
||||
from loguru import logger
|
||||
|
||||
return logger
|
||||
|
|
|
|||
|
|
@ -119,7 +119,7 @@ update_user_message_suffix: |
|
|||
Keep each section concise. Preserve exact file paths, function names, and error messages.
|
||||
|
||||
update_user_message_prefix_zh: |
|
||||
上述消息是要整合到现有摘要中的新对话消息,这些消息在<previous-summary>标签中提供。
|
||||
以上消息是需要整合到现有摘要中的新对话内容,现有摘要位于<previous-summary>标签中。
|
||||
|
||||
update_user_message_suffix_zh: |
|
||||
用新信息更新现有的结构化摘要。规则:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue