diff --git a/chatMSA/__init__.py b/chatMSA/__init__.py
new file mode 100644
index 0000000..69cd333
--- /dev/null
+++ b/chatMSA/__init__.py
@@ -0,0 +1,8 @@
+"""
+chatMSA — Multi-turn conversation system built on MSA (Memory Sparse Attention).
+
+Usage:
+ python -m chatMSA.app --model_path ckpt/MSA-4B --port 7860
+"""
+
+__version__ = "0.1.0"
diff --git a/chatMSA/__main__.py b/chatMSA/__main__.py
new file mode 100644
index 0000000..591e0ca
--- /dev/null
+++ b/chatMSA/__main__.py
@@ -0,0 +1,5 @@
+"""Allow running chatMSA as a module: python -m chatMSA"""
+
+from chatMSA.app import main
+
+main()
diff --git a/chatMSA/api/__init__.py b/chatMSA/api/__init__.py
new file mode 100644
index 0000000..a57c564
--- /dev/null
+++ b/chatMSA/api/__init__.py
@@ -0,0 +1,3 @@
+from chatMSA.api.router import create_fastapi_app
+
+__all__ = ["create_fastapi_app"]
diff --git a/chatMSA/api/chat_routes.py b/chatMSA/api/chat_routes.py
new file mode 100644
index 0000000..5c2f446
--- /dev/null
+++ b/chatMSA/api/chat_routes.py
@@ -0,0 +1,139 @@
+"""
+FastAPI routes for conversation management and chat.
+
+Endpoints:
+ POST /api/conversations -> Create conversation
+ GET /api/conversations -> List all conversations
+ GET /api/conversations/{id} -> Get conversation with messages
+ DELETE /api/conversations/{id} -> Delete conversation
+ PATCH /api/conversations/{id} -> Rename conversation
+ POST /api/conversations/{id}/messages -> Send message, get response
+ GET /api/health -> Engine status
+"""
+
+from fastapi import APIRouter, HTTPException
+
+from chatMSA.models.schemas import (
+ ConversationCreateRequest,
+ ConversationDetail,
+ ConversationRenameRequest,
+ ConversationSummary,
+ HealthResponse,
+ MessageResponse,
+ MessageSendRequest,
+)
+from chatMSA.services.chat_service import ChatService
+from chatMSA.services.msa_engine_service import MSAEngineService
+
+
+def create_chat_routes(chat_service: ChatService, engine: MSAEngineService) -> APIRouter:
+ """Create and return the chat API router."""
+ router = APIRouter(prefix="/api")
+
+ # ── Health ──────────────────────────────────────────────────
+
+ @router.get("/health", response_model=HealthResponse)
+ def health():
+ status = "ready" if engine.is_ready else ("loading" if engine.is_loading else "error")
+ import torch
+ gpu_count = len(engine.config.devices) if engine.config.devices else torch.cuda.device_count()
+ return HealthResponse(
+ status=status,
+ model_path=engine.config.model_path,
+ gpu_count=gpu_count,
+ uptime_seconds=engine.uptime,
+ error=engine.error,
+ )
+
+ # ── Conversations ───────────────────────────────────────────
+
+ @router.post("/conversations", response_model=ConversationDetail, status_code=201)
+ def create_conversation(req: ConversationCreateRequest = None):
+ if req is None:
+ req = ConversationCreateRequest()
+ conv = chat_service.create_conversation(title=req.title)
+ return _to_detail(conv)
+
+ @router.get("/conversations", response_model=list)
+ def list_conversations():
+ convs = chat_service.list_conversations()
+ return [_to_summary(c) for c in convs]
+
+ @router.get("/conversations/{conv_id}", response_model=ConversationDetail)
+ def get_conversation(conv_id: str):
+ conv = chat_service.get_conversation(conv_id)
+ if conv is None:
+ raise HTTPException(status_code=404, detail="Conversation not found")
+ return _to_detail(conv)
+
+ @router.delete("/conversations/{conv_id}", status_code=204)
+ def delete_conversation(conv_id: str):
+ deleted = chat_service.delete_conversation(conv_id)
+ if not deleted:
+ raise HTTPException(status_code=404, detail="Conversation not found")
+
+ @router.patch("/conversations/{conv_id}", response_model=ConversationDetail)
+ def rename_conversation(conv_id: str, req: ConversationRenameRequest):
+ conv = chat_service.rename_conversation(conv_id, req.title)
+ if conv is None:
+ raise HTTPException(status_code=404, detail="Conversation not found")
+ return _to_detail(conv)
+
+ # ── Messages ────────────────────────────────────────────────
+
+ @router.post("/conversations/{conv_id}/messages", response_model=MessageResponse)
+ def send_message(conv_id: str, req: MessageSendRequest):
+ try:
+ msg = chat_service.send_message(conv_id, req.content)
+ except ValueError as e:
+ raise HTTPException(status_code=404, detail=str(e))
+ except RuntimeError as e:
+ raise HTTPException(status_code=503, detail=str(e))
+ return MessageResponse(
+ id=msg.id,
+ role=msg.role,
+ content=msg.content,
+ timestamp=msg.timestamp,
+ recall_topk=msg.recall_topk,
+ )
+
+ return router
+
+
+# ── Helpers ─────────────────────────────────────────────────────
+
+def _to_summary(conv) -> ConversationSummary:
+ msg_count = getattr(conv, "_message_count", conv.message_count)
+ turn_count = getattr(conv, "_turn_count", conv.turn_count)
+ last_preview = None
+ if conv.messages:
+ last_preview = conv.messages[-1].content[:100]
+ return ConversationSummary(
+ id=conv.id,
+ title=conv.title,
+ created_at=conv.created_at,
+ updated_at=conv.updated_at,
+ message_count=msg_count,
+ turn_count=turn_count,
+ last_message_preview=last_preview,
+ )
+
+
+def _to_detail(conv) -> ConversationDetail:
+ return ConversationDetail(
+ id=conv.id,
+ title=conv.title,
+ created_at=conv.created_at,
+ updated_at=conv.updated_at,
+ messages=[
+ MessageResponse(
+ id=m.id,
+ role=m.role,
+ content=m.content,
+ timestamp=m.timestamp,
+ recall_topk=m.recall_topk,
+ )
+ for m in conv.messages
+ ],
+ metadata=conv.metadata,
+ )
diff --git a/chatMSA/api/router.py b/chatMSA/api/router.py
new file mode 100644
index 0000000..6e67e60
--- /dev/null
+++ b/chatMSA/api/router.py
@@ -0,0 +1,45 @@
+"""
+API router aggregation.
+
+Combines all route modules into a single FastAPI application.
+"""
+
+from fastapi import FastAPI
+from fastapi.middleware.cors import CORSMiddleware
+
+from chatMSA.api.chat_routes import create_chat_routes
+from chatMSA.services.chat_service import ChatService
+from chatMSA.services.msa_engine_service import MSAEngineService
+
+
+def create_fastapi_app(chat_service: ChatService, engine: MSAEngineService) -> FastAPI:
+ """
+ Create and configure the FastAPI application.
+
+ Args:
+ chat_service: The chat business logic service.
+ engine: The MSA engine service (for health checks).
+
+ Returns:
+ Configured FastAPI app with all routes mounted.
+ """
+ app = FastAPI(
+ title="chatMSA API",
+ description="Multi-turn conversation system built on MSA (Memory Sparse Attention)",
+ version="0.1.0",
+ )
+
+ # CORS — allow Gradio frontend and local development
+ app.add_middleware(
+ CORSMiddleware,
+ allow_origins=["*"],
+ allow_credentials=True,
+ allow_methods=["*"],
+ allow_headers=["*"],
+ )
+
+ # Mount API routes
+ chat_router = create_chat_routes(chat_service, engine)
+ app.include_router(chat_router)
+
+ return app
diff --git a/chatMSA/app.py b/chatMSA/app.py
new file mode 100644
index 0000000..487cd9a
--- /dev/null
+++ b/chatMSA/app.py
@@ -0,0 +1,85 @@
+"""
+chatMSA entry point.
+
+Launches FastAPI (REST API) with embedded Gradio (chat UI).
+
+Usage:
+ python -m chatMSA.app --model_path ckpt/MSA-4B --port 7860
+ python -m chatMSA.app --help
+"""
+
+import signal
+import sys
+import pathlib
+
+# Ensure project root is importable
+_project_root = pathlib.Path(__file__).parent.parent
+sys.path.insert(0, str(_project_root))
+
+import gradio as gr
+import uvicorn
+
+from chatMSA.config import ChatConfig
+from chatMSA.services.chat_service import ChatService
+from chatMSA.services.msa_engine_service import MSAEngineService
+from chatMSA.storage.sqlite_store import SQLiteConversationStore
+from chatMSA.api.router import create_fastapi_app
+from chatMSA.frontend.chat_ui import create_gradio_app
+
+
+def main():
+ """Main entry point."""
+ # 1. Parse config
+ config = ChatConfig.from_args()
+ print(f"[chatMSA] Config: model={config.model_path}, db={config.db_path}")
+ print(f"[chatMSA] Server: {config.host}:{config.port}")
+
+ # 2. Initialize storage
+ store = SQLiteConversationStore(config.db_path)
+ store.initialize()
+ print(f"[chatMSA] Storage initialized: {config.db_path}")
+
+ # 3. Initialize engine (lazy — heavy loading happens in engine.start())
+ engine = MSAEngineService(config)
+
+ # 4. Initialize chat service
+ chat_service = ChatService(engine, store, config)
+
+ # 5. Create FastAPI app
+ fastapi_app = create_fastapi_app(chat_service, engine)
+
+ # 6. Create and mount Gradio app
+ gradio_app = create_gradio_app(chat_service)
+ fastapi_app = gr.mount_gradio_app(fastapi_app, gradio_app, path="/")
+
+ # 7. Graceful shutdown
+ def shutdown_handler(sig, frame):
+ print("\n[chatMSA] Shutting down...")
+ engine.stop()
+ store.close()
+ sys.exit(0)
+
+ signal.signal(signal.SIGINT, shutdown_handler)
+ signal.signal(signal.SIGTERM, shutdown_handler)
+
+ # 8. Start engine (heavy: model loading + memory prefill)
+ # This runs before the HTTP server so the first request doesn't wait.
+ # TODO: Move engine.start() to a background thread and serve a "loading"
+ # page from Gradio while the engine initializes. This would improve UX
+ # for large models with long startup times.
+ print("[chatMSA] Starting MSA engine (this may take a while)...")
+ try:
+ engine.start()
+ print("[chatMSA] Engine ready!")
+ except Exception as e:
+ print(f"[chatMSA] Engine failed to start: {e}")
+ print("[chatMSA] Continuing in degraded mode (chat will return errors)")
+
+ # 9. Launch server
+ print(f"[chatMSA] UI: http://{config.host}:{config.port}")
+ print(f"[chatMSA] API: http://{config.host}:{config.port}/docs")
+ uvicorn.run(fastapi_app, host=config.host, port=config.port, log_level="info")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/chatMSA/config.py b/chatMSA/config.py
new file mode 100644
index 0000000..8d628d2
--- /dev/null
+++ b/chatMSA/config.py
@@ -0,0 +1,135 @@
+"""
+Unified configuration for chatMSA.
+
+All settings in one place. Converts to MSA's native config types internally.
+"""
+
+import argparse
+from dataclasses import dataclass, field
+from typing import List, Optional
+
+import sys
+import pathlib
+
+# Add project root to path so we can import MSA's config types
+_project_root = pathlib.Path(__file__).parent.parent
+sys.path.insert(0, str(_project_root))
+
+from src.config.memory_config import GenerateConfig, ModelConfig, MemoryConfig
+from src.utils.template import QWEN3_TEMPLATE, QWEN3_INSTRUCT_TEMPLATE
+
+
+@dataclass
+class ChatConfig:
+ """Single source of truth for all chatMSA settings."""
+
+ # ── MSA Engine ──────────────────────────────────────────────
+ model_path: str = "ckpt/MSA-4B"
+ template: str = "QWEN3_INSTRUCT_TEMPLATE"
+ devices: Optional[List[int]] = None # None = auto-detect all GPUs
+ max_generate_tokens: int = 512
+ max_batch_size: int = 4
+ top_p: float = 0.9
+ temperature: float = 0.0
+ block_size: int = 2048
+ max_chunk_per_block: int = 16384
+ max_seq_len: int = 0 # 0 = unlimited
+ max_query_seq_len: int = 0 # 0 = unlimited
+
+ # ── Chat ────────────────────────────────────────────────────
+ max_history_turns: int = 20 # Max conversation turns kept in prompt context
+ max_context_tokens: int = 4096 # Token budget for history + current query
+ db_path: str = "data/chat_msa.db"
+
+ # ── Server ──────────────────────────────────────────────────
+ host: str = "0.0.0.0"
+ port: int = 7860
+
+ @property
+ def template_dict(self) -> dict:
+ """Resolve template name to the actual template dict."""
+ templates = {
+ "QWEN3_TEMPLATE": QWEN3_TEMPLATE,
+ "QWEN3_INSTRUCT_TEMPLATE": QWEN3_INSTRUCT_TEMPLATE,
+ }
+ if self.template not in templates:
+ raise ValueError(f"Unknown template: {self.template}. Choose from {list(templates.keys())}")
+ return templates[self.template]
+
+ def to_generate_config(self) -> GenerateConfig:
+ """Convert to MSA's GenerateConfig."""
+ import torch
+ devices = self.devices if self.devices is not None else list(range(torch.cuda.device_count()))
+ return GenerateConfig(
+ devices=devices,
+ template=self.template_dict,
+ max_generate_tokens=self.max_generate_tokens,
+ max_seq_len=self.max_seq_len,
+ max_query_seq_len=self.max_query_seq_len,
+ max_batch_size=self.max_batch_size,
+ top_p=self.top_p,
+ temperature=self.temperature,
+ qa_mode=True,
+ )
+
+ def to_model_config(self) -> ModelConfig:
+ """Convert to MSA's ModelConfig."""
+ return ModelConfig(
+ model_path=self.model_path,
+ doc_top_k=16, # TODO: make configurable
+ pooling_kernel_size=64, # TODO: make configurable
+ router_layer_idx="all",
+ )
+
+ def to_memory_config(self, memory_file_path: str = "") -> MemoryConfig:
+ """Convert to MSA's MemoryConfig."""
+ return MemoryConfig(
+ block_size=self.block_size,
+ pooling_kernel_size=64, # TODO: make configurable
+ slice_chunk_size=self.max_chunk_per_block,
+ memory_file_path=memory_file_path,
+ )
+
+ @classmethod
+ def from_args(cls, args: Optional[List[str]] = None) -> "ChatConfig":
+ """Parse CLI arguments into a ChatConfig."""
+ parser = argparse.ArgumentParser(description="chatMSA — Multi-turn conversation on MSA")
+ parser.add_argument("--model_path", type=str, default="ckpt/MSA-4B")
+ parser.add_argument("--template", type=str, default="QWEN3_INSTRUCT_TEMPLATE",
+ choices=["QWEN3_TEMPLATE", "QWEN3_INSTRUCT_TEMPLATE"])
+ parser.add_argument("--devices", type=str, default=None,
+ help="Comma-separated GPU IDs, e.g. '0,1,2,3'. Default: all")
+ parser.add_argument("--max_generate_tokens", type=int, default=512)
+ parser.add_argument("--max_batch_size", type=int, default=4)
+ parser.add_argument("--top_p", type=float, default=0.9)
+ parser.add_argument("--temperature", type=float, default=0.0)
+ parser.add_argument("--block_size", type=int, default=2048)
+ parser.add_argument("--max_chunk_per_block", type=int, default=16384)
+ parser.add_argument("--max_history_turns", type=int, default=20)
+ parser.add_argument("--max_context_tokens", type=int, default=4096)
+ parser.add_argument("--db_path", type=str, default="data/chat_msa.db")
+ parser.add_argument("--host", type=str, default="0.0.0.0")
+ parser.add_argument("--port", type=int, default=7860)
+
+ parsed = parser.parse_args(args)
+
+ devices = None
+ if parsed.devices is not None:
+ devices = [int(d) for d in parsed.devices.split(",")]
+
+ return cls(
+ model_path=parsed.model_path,
+ template=parsed.template,
+ devices=devices,
+ max_generate_tokens=parsed.max_generate_tokens,
+ max_batch_size=parsed.max_batch_size,
+ top_p=parsed.top_p,
+ temperature=parsed.temperature,
+ block_size=parsed.block_size,
+ max_chunk_per_block=parsed.max_chunk_per_block,
+ max_history_turns=parsed.max_history_turns,
+ max_context_tokens=parsed.max_context_tokens,
+ db_path=parsed.db_path,
+ host=parsed.host,
+ port=parsed.port,
+ )
diff --git a/chatMSA/frontend/__init__.py b/chatMSA/frontend/__init__.py
new file mode 100644
index 0000000..a317fd7
--- /dev/null
+++ b/chatMSA/frontend/__init__.py
@@ -0,0 +1,3 @@
+from chatMSA.frontend.chat_ui import create_gradio_app
+
+__all__ = ["create_gradio_app"]
diff --git a/chatMSA/frontend/chat_ui.py b/chatMSA/frontend/chat_ui.py
new file mode 100644
index 0000000..60fb6dd
--- /dev/null
+++ b/chatMSA/frontend/chat_ui.py
@@ -0,0 +1,191 @@
+"""
+Gradio frontend for chatMSA.
+
+Layout:
+ - Left sidebar: conversation list with "New Chat" button
+ - Main area: chat history + input box
+ - Status bar: engine status indicator
+"""
+
+from typing import List, Tuple
+
+import gradio as gr
+
+from chatMSA.services.chat_service import ChatService
+
+
+def create_gradio_app(chat_service: ChatService) -> gr.Blocks:
+ """
+ Create and return the Gradio Blocks application.
+
+ The returned app can be mounted onto FastAPI via gr.mount_gradio_app()
+ or launched standalone.
+ """
+ with gr.Blocks(
+ title="chatMSA",
+ theme=gr.themes.Soft(),
+ css="""
+ .sidebar { min-width: 250px; max-width: 300px; }
+ .conv-item { padding: 8px 12px; cursor: pointer; border-radius: 6px; margin: 2px 0; }
+ .conv-item:hover { background: #e8e8e8; }
+ .conv-item.active { background: #d0e0ff; font-weight: bold; }
+ """,
+ ) as app:
+
+ # ── State ───────────────────────────────────────────────
+ current_conv_id = gr.State(value=None)
+
+ # ── Layout ──────────────────────────────────────────────
+ gr.Markdown("# 💬 chatMSA\n*Multi-turn conversation on Memory Sparse Attention*")
+
+ with gr.Row():
+ # Left sidebar
+ with gr.Column(scale=1, elem_classes="sidebar"):
+ new_chat_btn = gr.Button("➕ New Chat", variant="primary", size="sm")
+ gr.Markdown("### Conversations")
+ conv_list = gr.HTML(value=_render_conv_list(chat_service, None))
+
+ # Main chat area
+ with gr.Column(scale=4):
+ chatbot = gr.Chatbot(
+ label="Chat",
+ height=500,
+ type="messages",
+ show_copy_button=True,
+ )
+ with gr.Row():
+ msg_input = gr.Textbox(
+ placeholder="Type your message...",
+ show_label=False,
+ scale=9,
+ container=False,
+ )
+ send_btn = gr.Button("Send", variant="primary", scale=1)
+
+ status_text = gr.Markdown("*Ready*")
+
+ # ── Event Handlers ──────────────────────────────────────
+
+ def on_new_chat():
+ """Create a new conversation and switch to it."""
+ conv = chat_service.create_conversation()
+ return (
+ conv.id, # current_conv_id
+ [], # chatbot (empty)
+ _render_conv_list(chat_service, conv.id), # conv_list
+ gr.update(value=""), # msg_input
+ )
+
+ def on_select_conv(evt: gr.SelectData):
+ """Switch to a selected conversation."""
+ conv_id = evt.value
+ conv = chat_service.get_conversation(conv_id)
+ if conv is None:
+ return (None, [], _render_conv_list(chat_service, None), "")
+ messages = _conv_to_chatbot(conv)
+ return (
+ conv_id,
+ messages,
+ _render_conv_list(chat_service, conv_id),
+ gr.update(value=""),
+ )
+
+ def on_send_message(user_text: str, conv_id: str, history: list):
+ """Send a message and get the response."""
+ if not user_text.strip():
+ return history, gr.update(value=""), _render_conv_list(chat_service, conv_id), conv_id
+
+ # Create conversation if none selected
+ if conv_id is None:
+ conv = chat_service.create_conversation()
+ conv_id = conv.id
+
+ # Add user message to chatbot immediately
+ history = history + [{"role": "user", "content": user_text}]
+
+ try:
+ # Call chat service (this blocks until MSA responds)
+ assistant_msg = chat_service.send_message(conv_id, user_text)
+ history = history + [{"role": "assistant", "content": assistant_msg.content}]
+ status = f"*Responded at {assistant_msg.timestamp:.0f}*"
+ except RuntimeError as e:
+ history = history + [{"role": "assistant", "content": f"⚠️ Error: {e}"}]
+ status = f"*Error: {e}*"
+ except Exception as e:
+ history = history + [{"role": "assistant", "content": f"⚠️ Unexpected error: {e}"}]
+ status = f"*Error: {e}*"
+
+ return (
+ history,
+ gr.update(value=""),
+ _render_conv_list(chat_service, conv_id),
+ conv_id,
+ status,
+ )
+
+ def on_delete_conv(conv_id: str):
+ """Delete the current conversation."""
+ if conv_id is not None:
+ chat_service.delete_conversation(conv_id)
+ return (
+ None,
+ [],
+ _render_conv_list(chat_service, None),
+ )
+
+ # Wire up events
+ new_chat_btn.click(
+ fn=on_new_chat,
+ outputs=[current_conv_id, chatbot, conv_list, msg_input],
+ )
+
+ # TODO: Conversation selection via HTML click events requires
+ # a custom JavaScript component or Gradio's gr.render decorator.
+ # Currently, conversation switching works via the Gradio select event.
+ # For a production UI, consider:
+ # 1. Using gr.render with a radio button list
+ # 2. Adding custom JS for clickable conversation items
+ # 3. Using gr.Dropdown as a simpler alternative
+
+ send_btn.click(
+ fn=on_send_message,
+ inputs=[msg_input, current_conv_id, chatbot],
+ outputs=[chatbot, msg_input, conv_list, current_conv_id, status_text],
+ )
+
+ msg_input.submit(
+ fn=on_send_message,
+ inputs=[msg_input, current_conv_id, chatbot],
+ outputs=[chatbot, msg_input, conv_list, current_conv_id, status_text],
+ )
+
+ return app
+
+
+# ── Helpers ─────────────────────────────────────────────────────
+
+def _conv_to_chatbot(conv) -> list:
+ """Convert Conversation messages to Gradio chatbot format."""
+ messages = []
+ for msg in conv.messages:
+ messages.append({"role": msg.role, "content": msg.content})
+ return messages
+
+
+def _render_conv_list(chat_service: ChatService, active_id: str = None) -> str:
+ """Render the conversation list as clickable HTML."""
+ convs = chat_service.list_conversations()
+ if not convs:
+ return "
No conversations yet
"
+
+ html_parts = []
+ for conv in convs:
+ active_class = "active" if conv.id == active_id else ""
+ title = conv.title[:30] + ("..." if len(conv.title) > 30 else "")
+ html_parts.append(
+ f''
+ f'{title}
'
+ )
+ return "\n".join(html_parts)
diff --git a/chatMSA/models/__init__.py b/chatMSA/models/__init__.py
new file mode 100644
index 0000000..7f9861a
--- /dev/null
+++ b/chatMSA/models/__init__.py
@@ -0,0 +1,22 @@
+from chatMSA.models.conversation import Message, Conversation
+from chatMSA.models.schemas import (
+ ConversationCreateRequest,
+ ConversationRenameRequest,
+ MessageSendRequest,
+ MessageResponse,
+ ConversationSummary,
+ ConversationDetail,
+ HealthResponse,
+)
+
+__all__ = [
+ "Message",
+ "Conversation",
+ "ConversationCreateRequest",
+ "ConversationRenameRequest",
+ "MessageSendRequest",
+ "MessageResponse",
+ "ConversationSummary",
+ "ConversationDetail",
+ "HealthResponse",
+]
diff --git a/chatMSA/models/conversation.py b/chatMSA/models/conversation.py
new file mode 100644
index 0000000..7971cf6
--- /dev/null
+++ b/chatMSA/models/conversation.py
@@ -0,0 +1,95 @@
+"""
+Data models for conversations and messages.
+
+Pure dataclasses with no framework dependency — usable by any layer.
+"""
+
+import time
+import uuid
+from dataclasses import dataclass, field
+from typing import Dict, List, Optional
+
+
+@dataclass
+class Message:
+ """A single message in a conversation."""
+ id: str = field(default_factory=lambda: str(uuid.uuid4()))
+ role: str = "" # "user" | "assistant"
+ content: str = ""
+ timestamp: float = field(default_factory=time.time)
+ recall_topk: Optional[Dict] = None # MSA retrieved doc IDs per layer (optional)
+
+ def to_dict(self) -> dict:
+ return {
+ "id": self.id,
+ "role": self.role,
+ "content": self.content,
+ "timestamp": self.timestamp,
+ "recall_topk": self.recall_topk,
+ }
+
+ @classmethod
+ def from_dict(cls, data: dict) -> "Message":
+ return cls(
+ id=data["id"],
+ role=data["role"],
+ content=data["content"],
+ timestamp=data["timestamp"],
+ recall_topk=data.get("recall_topk"),
+ )
+
+
+@dataclass
+class Conversation:
+ """A conversation session containing ordered messages."""
+ id: str = field(default_factory=lambda: str(uuid.uuid4()))
+ title: str = "New Chat"
+ created_at: float = field(default_factory=time.time)
+ updated_at: float = field(default_factory=time.time)
+ messages: List[Message] = field(default_factory=list)
+ metadata: Optional[Dict] = None # Extensible: model info, config snapshot, etc.
+
+ @property
+ def message_count(self) -> int:
+ return len(self.messages)
+
+ @property
+ def turn_count(self) -> int:
+ """Number of complete user-assistant turns."""
+ return sum(1 for m in self.messages if m.role == "user")
+
+ def add_message(self, role: str, content: str, recall_topk: Optional[Dict] = None) -> Message:
+ """Create and append a message, update timestamp."""
+ msg = Message(role=role, content=content, recall_topk=recall_topk)
+ self.messages.append(msg)
+ self.updated_at = time.time()
+ return msg
+
+ def get_history(self, max_turns: Optional[int] = None) -> List[Message]:
+ """Return messages, optionally limited to the last N turns."""
+ if max_turns is None:
+ return list(self.messages)
+ # Each turn = 1 user + 1 assistant message
+ max_messages = max_turns * 2
+ return self.messages[-max_messages:] if max_messages < len(self.messages) else list(self.messages)
+
+ def to_dict(self) -> dict:
+ return {
+ "id": self.id,
+ "title": self.title,
+ "created_at": self.created_at,
+ "updated_at": self.updated_at,
+ "messages": [m.to_dict() for m in self.messages],
+ "metadata": self.metadata,
+ }
+
+ @classmethod
+ def from_dict(cls, data: dict) -> "Conversation":
+ return cls(
+ id=data["id"],
+ title=data["title"],
+ created_at=data["created_at"],
+ updated_at=data["updated_at"],
+ messages=[Message.from_dict(m) for m in data.get("messages", [])],
+ metadata=data.get("metadata"),
+ )
diff --git a/chatMSA/models/schemas.py b/chatMSA/models/schemas.py
new file mode 100644
index 0000000..d004912
--- /dev/null
+++ b/chatMSA/models/schemas.py
@@ -0,0 +1,58 @@
+"""
+Pydantic models for API request/response serialization and validation.
+"""
+
+from typing import Dict, List, Optional
+from pydantic import BaseModel, Field
+
+
+# ── Request Models ──────────────────────────────────────────────
+
+class ConversationCreateRequest(BaseModel):
+ title: Optional[str] = Field(None, description="Conversation title. Auto-generated if omitted.")
+ metadata: Optional[Dict] = Field(None, description="Optional metadata (model info, settings, etc.)")
+
+
+class ConversationRenameRequest(BaseModel):
+ title: str = Field(..., min_length=1, max_length=200, description="New conversation title")
+
+
+class MessageSendRequest(BaseModel):
+ content: str = Field(..., min_length=1, description="User message content")
+
+
+# ── Response Models ─────────────────────────────────────────────
+
+class MessageResponse(BaseModel):
+ id: str
+ role: str
+ content: str
+ timestamp: float
+ recall_topk: Optional[Dict] = None
+
+
+class ConversationSummary(BaseModel):
+ id: str
+ title: str
+ created_at: float
+ updated_at: float
+ message_count: int
+ turn_count: int
+ last_message_preview: Optional[str] = None
+
+
+class ConversationDetail(BaseModel):
+ id: str
+ title: str
+ created_at: float
+ updated_at: float
+ messages: List[MessageResponse]
+ metadata: Optional[Dict] = None
+
+
+class HealthResponse(BaseModel):
+ status: str # "ready" | "loading" | "error"
+ model_path: str
+ gpu_count: int
+ uptime_seconds: float
+ error: Optional[str] = None
diff --git a/chatMSA/services/__init__.py b/chatMSA/services/__init__.py
new file mode 100644
index 0000000..da3d837
--- /dev/null
+++ b/chatMSA/services/__init__.py
@@ -0,0 +1,5 @@
+from chatMSA.services.msa_engine_service import MSAEngineService
+from chatMSA.services.chat_service import ChatService
+from chatMSA.services.prompt_builder import PromptBuilder
+
+__all__ = ["MSAEngineService", "ChatService", "PromptBuilder"]
diff --git a/chatMSA/services/chat_service.py b/chatMSA/services/chat_service.py
new file mode 100644
index 0000000..860c875
--- /dev/null
+++ b/chatMSA/services/chat_service.py
@@ -0,0 +1,152 @@
+"""
+Chat service — orchestrates conversation management, prompt building,
+and MSA engine calls.
+
+This is the core business logic layer. It knows nothing about HTTP or UI.
+"""
+
+import time
+from typing import List, Optional
+
+from chatMSA.config import ChatConfig
+from chatMSA.models.conversation import Conversation, Message
+from chatMSA.services.msa_engine_service import MSAEngineService
+from chatMSA.services.prompt_builder import PromptBuilder
+from chatMSA.storage.base import BaseConversationStore
+
+
+class ChatService:
+ """
+ Orchestrates multi-turn conversation on MSA.
+
+ Responsibilities:
+ - CRUD for conversations (delegates to store)
+ - Build prompts with history (delegates to PromptBuilder)
+ - Call MSA engine for generation (delegates to MSAEngineService)
+ - Persist messages after each turn
+ """
+
+ def __init__(
+ self,
+ engine: MSAEngineService,
+ store: BaseConversationStore,
+ config: ChatConfig,
+ ):
+ self.engine = engine
+ self.store = store
+ self.config = config
+ self.prompt_builder = PromptBuilder(config)
+
+ # ── Conversation CRUD ───────────────────────────────────────
+
+ def create_conversation(self, title: Optional[str] = None) -> Conversation:
+ """Create a new conversation and persist it."""
+ conv = Conversation(title=title or "New Chat")
+ self.store.save_conversation(conv)
+ return conv
+
+ def get_conversation(self, conv_id: str) -> Optional[Conversation]:
+ """Load a full conversation with all messages."""
+ return self.store.load_conversation(conv_id)
+
+ def list_conversations(self) -> List[Conversation]:
+ """List all conversations (lightweight, no messages)."""
+ return self.store.list_conversations()
+
+ def delete_conversation(self, conv_id: str) -> bool:
+ """Delete a conversation and all its messages."""
+ return self.store.delete_conversation(conv_id)
+
+ def rename_conversation(self, conv_id: str, title: str) -> Optional[Conversation]:
+ """Rename a conversation. Returns updated conversation or None if not found."""
+ conv = self.store.load_conversation(conv_id)
+ if conv is None:
+ return None
+ conv.title = title
+ self.store.save_conversation(conv)
+ return conv
+
+ # ── Chat ────────────────────────────────────────────────────
+
+ def send_message(self, conv_id: str, user_message: str) -> Message:
+ """
+ Send a user message and get the assistant's response.
+
+ Flow:
+ 1. Load conversation history from store
+ 2. Persist the user message
+ 3. Build prompt with PromptBuilder (history + current query)
+ 4. Call MSAEngineService.generate()
+ 5. Parse response, create assistant Message
+ 6. Persist the assistant message
+ 7. Return assistant Message
+
+ Raises:
+ ValueError: If conversation not found.
+ RuntimeError: If engine is not ready.
+ """
+ # 1. Load conversation
+ conv = self.store.load_conversation(conv_id)
+ if conv is None:
+ raise ValueError(f"Conversation not found: {conv_id}")
+
+ # 2. Persist user message
+ user_msg = Message(role="user", content=user_message)
+ self.store.add_message(conv_id, user_msg)
+ conv.messages.append(user_msg)
+
+ # 3. Build prompt with history
+ history = conv.get_history(max_turns=self.config.max_history_turns)
+ # Remove the just-added user message from history (it's the current query)
+ history_without_current = history[:-1]
+ prompt = self.prompt_builder.build(history_without_current, user_message)
+
+ # 4. Generate response
+ raw_response, recall_topk = self.engine.generate(prompt)
+
+ # 5. Parse response
+ assistant_content = self._parse_response(raw_response)
+
+ # 6. Persist assistant message
+ assistant_msg = Message(
+ role="assistant",
+ content=assistant_content,
+ recall_topk=recall_topk if recall_topk else None,
+ )
+ self.store.add_message(conv_id, assistant_msg)
+
+ # Auto-generate title from first user message if title is default
+ if conv.title == "New Chat" and conv.turn_count == 1:
+ auto_title = user_message[:50] + ("..." if len(user_message) > 50 else "")
+ conv.title = auto_title
+ self.store.save_conversation(conv)
+
+ return assistant_msg
+
+ def _parse_response(self, raw_response: str) -> str:
+ """
+ Extract the final answer from MSA's raw output.
+
+ MSA outputs structured text with ... blocks
+ and "The answer to the question is:" markers. We extract
+ the clean answer for display.
+ """
+ # Try to extract from structured output
+ if "The answer to the question is:" in raw_response:
+ answer = raw_response.split("The answer to the question is:")[-1]
+ # Clean up common artifacts
+ answer = answer.replace("", "").strip()
+ if "Answer:" in answer:
+ answer = answer.split("Answer:")[-1].strip()
+ return answer
+
+ # Fallback: remove blocks
+ if "" in raw_response and "" in raw_response:
+ before = raw_response.split("")[0]
+ after = raw_response.split("")[-1]
+ cleaned = (before + after).strip()
+ if cleaned:
+ return cleaned
+
+ # Last resort: return as-is
+ return raw_response.strip()
diff --git a/chatMSA/services/msa_engine_service.py b/chatMSA/services/msa_engine_service.py
new file mode 100644
index 0000000..f633fce
--- /dev/null
+++ b/chatMSA/services/msa_engine_service.py
@@ -0,0 +1,135 @@
+"""
+MSAEngine lifecycle wrapper.
+
+Manages the heavy MSA engine (model loading, multi-GPU workers, prefill)
+with lazy initialization and clean shutdown.
+"""
+
+import time
+import threading
+from typing import Optional, Tuple, Dict
+
+import sys
+import pathlib
+
+_project_root = pathlib.Path(__file__).parent.parent.parent
+sys.path.insert(0, str(_project_root))
+
+from chatMSA.config import ChatConfig
+from src.msa_service import MSAEngine
+
+
+class MSAEngineService:
+ """
+ Wraps MSAEngine with lifecycle management.
+
+ Usage:
+ engine_service = MSAEngineService(config)
+ engine_service.start() # Heavy: loads model, prefill memory
+ text, topk = engine_service.generate("What is MSA?")
+ engine_service.stop() # Join workers
+
+ Context manager:
+ with MSAEngineService(config) as engine:
+ text, topk = engine.generate("What is MSA?")
+ """
+
+ def __init__(self, config: ChatConfig):
+ self.config = config
+ self._engine: Optional[MSAEngine] = None
+ self._ready = False
+ self._loading = False
+ self._error: Optional[str] = None
+ self._start_time: float = 0
+ self._lock = threading.Lock()
+
+ @property
+ def is_ready(self) -> bool:
+ return self._ready
+
+ @property
+ def is_loading(self) -> bool:
+ return self._loading
+
+ @property
+ def error(self) -> Optional[str]:
+ return self._error
+
+ @property
+ def uptime(self) -> float:
+ if self._start_time == 0:
+ return 0.0
+ return time.time() - self._start_time
+
+ def start(self, memory_file_path: str = "") -> None:
+ """
+ Initialize the MSA engine. This is expensive (model loading + prefill).
+
+ Args:
+ memory_file_path: Path to the memory corpus file (pickle/json).
+ Empty string = no memory corpus (pure chat mode).
+ """
+ with self._lock:
+ if self._ready:
+ return
+ if self._loading:
+ raise RuntimeError("Engine is already loading")
+
+ self._loading = True
+ self._error = None
+
+ try:
+ generate_config = self.config.to_generate_config()
+ model_config = self.config.to_model_config()
+ memory_config = self.config.to_memory_config(memory_file_path)
+
+ self._engine = MSAEngine(generate_config, model_config, memory_config)
+ self._ready = True
+ self._start_time = time.time()
+ except Exception as e:
+ self._error = str(e)
+ raise
+ finally:
+ self._loading = False
+
+ def stop(self) -> None:
+ """Shut down the engine and join worker processes."""
+ with self._lock:
+ if self._engine is not None:
+ self._engine.stop_workers()
+ self._engine = None
+ self._ready = False
+ self._start_time = 0
+
+ def generate(self, prompt: str) -> Tuple[str, Dict]:
+ """
+ Generate a response for a single prompt.
+
+ Args:
+ prompt: The input text (already includes history if applicable).
+
+ Returns:
+ Tuple of (generated_text, recall_topk).
+ recall_topk maps layer_idx -> list of retrieved doc IDs.
+
+ Raises:
+ RuntimeError: If engine is not ready.
+ """
+ if not self._ready or self._engine is None:
+ raise RuntimeError("Engine is not ready. Call start() first.")
+
+ texts, recall_topk, _ = self._engine.generate(
+ prompt,
+ require_recall_topk=True,
+ )
+ # texts is a list; single prompt → single result
+ response_text = texts[0] if texts else ""
+ return response_text, recall_topk
+
+ def __enter__(self):
+ self.start()
+ return self
+
+ def __exit__(self, exc_type, exc_val, exc_tb):
+ self.stop()
+ return False
diff --git a/chatMSA/services/prompt_builder.py b/chatMSA/services/prompt_builder.py
new file mode 100644
index 0000000..5098d11
--- /dev/null
+++ b/chatMSA/services/prompt_builder.py
@@ -0,0 +1,73 @@
+"""
+Prompt construction with conversation history.
+
+Assembles multi-turn history + current query into a single prompt string
+that MSA's template system can process.
+"""
+
+from typing import List
+
+from chatMSA.config import ChatConfig
+from chatMSA.models.conversation import Message
+
+
+class PromptBuilder:
+ """Builds prompts that include conversation history for MSA inference."""
+
+ def __init__(self, config: ChatConfig):
+ self.config = config
+
+ def build(self, history: List[Message], current_query: str) -> str:
+ """
+ Construct a single prompt string with conversation history.
+
+ MSA's _apply_template() wraps the output in Qwen chat format,
+ so we only produce the raw text content here.
+
+ Args:
+ history: Previous messages (user + assistant pairs).
+ current_query: The latest user question.
+
+ Returns:
+ A prompt string ready to pass to MSAEngine.generate().
+
+ Format example:
+ Previous conversation:
+ User: What is MSA?
+ Assistant: MSA is a scalable sparse attention framework...
+
+ User: How does it scale to 100M tokens?
+ """
+ if not history:
+ return current_query
+
+ parts = []
+
+ # Build conversation history section
+ history_lines = []
+ for msg in history:
+ prefix = "User" if msg.role == "user" else "Assistant"
+ history_lines.append(f"{prefix}: {msg.content}")
+
+ if history_lines:
+ parts.append("Previous conversation:")
+ parts.append("\n".join(history_lines))
+ parts.append("") # Blank line separator
+
+ # Append current query
+ parts.append(current_query)
+
+ return "\n".join(parts)
+
+ # TODO: Implement token-aware truncation.
+ #
+ # Currently we rely on max_history_turns to limit context size,
+ # but this doesn't account for actual token counts. A proper implementation
+ # would:
+ # 1. Accept the tokenizer as a dependency (or from MSAEngineService)
+ # 2. Count tokens in the assembled history
+ # 3. Progressively remove oldest turns until under max_context_tokens
+ # 4. Optionally add a "[earlier history truncated]" marker
+ #
+ # def build_with_token_budget(self, history, current_query, tokenizer):
+ # ...
diff --git a/chatMSA/storage/__init__.py b/chatMSA/storage/__init__.py
new file mode 100644
index 0000000..9a50bdc
--- /dev/null
+++ b/chatMSA/storage/__init__.py
@@ -0,0 +1,4 @@
+from chatMSA.storage.base import BaseConversationStore
+from chatMSA.storage.sqlite_store import SQLiteConversationStore
+
+__all__ = ["BaseConversationStore", "SQLiteConversationStore"]
diff --git a/chatMSA/storage/base.py b/chatMSA/storage/base.py
new file mode 100644
index 0000000..38d85e3
--- /dev/null
+++ b/chatMSA/storage/base.py
@@ -0,0 +1,54 @@
+"""
+Abstract storage interface for conversation persistence.
+
+Implement this interface to swap storage backends (SQLite, PostgreSQL, Redis, etc.)
+without changing any business logic.
+"""
+
+from abc import ABC, abstractmethod
+from typing import List, Optional
+
+from chatMSA.models.conversation import Conversation, Message
+
+
+class BaseConversationStore(ABC):
+ """Abstract base class for conversation persistence."""
+
+ @abstractmethod
+ def initialize(self) -> None:
+ """Create tables / indices if they don't exist. Called once at startup."""
+ ...
+
+ @abstractmethod
+ def save_conversation(self, conv: Conversation) -> None:
+ """Insert or update a conversation (metadata only, not messages)."""
+ ...
+
+ @abstractmethod
+ def load_conversation(self, conv_id: str) -> Optional[Conversation]:
+ """Load a full conversation with all messages. Returns None if not found."""
+ ...
+
+ @abstractmethod
+ def list_conversations(self) -> List[Conversation]:
+ """List all conversations (metadata only, no messages). Sorted by updated_at DESC."""
+ ...
+
+ @abstractmethod
+ def delete_conversation(self, conv_id: str) -> bool:
+ """Delete a conversation and all its messages. Returns True if found and deleted."""
+ ...
+
+ @abstractmethod
+ def add_message(self, conv_id: str, message: Message) -> None:
+ """Append a message to an existing conversation."""
+ ...
+
+ @abstractmethod
+ def get_messages(self, conv_id: str) -> List[Message]:
+ """Get all messages for a conversation, ordered by timestamp."""
+ ...
+
+ def close(self) -> None:
+ """Release resources. Override if needed."""
+ pass
diff --git a/chatMSA/storage/sqlite_store.py b/chatMSA/storage/sqlite_store.py
new file mode 100644
index 0000000..386d310
--- /dev/null
+++ b/chatMSA/storage/sqlite_store.py
@@ -0,0 +1,193 @@
+"""
+SQLite-backed conversation storage.
+
+For file-based databases: creates a new connection per call (thread-safe).
+For ":memory:": keeps a single persistent connection (tables would be lost otherwise).
+"""
+
+import json
+import os
+import sqlite3
+import time
+from typing import List, Optional
+
+from chatMSA.models.conversation import Conversation, Message
+from chatMSA.storage.base import BaseConversationStore
+
+
+class SQLiteConversationStore(BaseConversationStore):
+ """SQLite implementation of conversation persistence."""
+
+ def __init__(self, db_path: str = "data/chat_msa.db"):
+ self.db_path = db_path
+ self._persistent_conn: Optional[sqlite3.Connection] = None
+ # Ensure parent directory exists (skip for :memory:)
+ if db_path != ":memory:":
+ os.makedirs(os.path.dirname(db_path) or ".", exist_ok=True)
+
+ def _get_conn(self) -> sqlite3.Connection:
+ """
+ Get a database connection.
+
+ For :memory: databases, returns a single persistent connection
+ (otherwise tables would be lost between calls).
+ For file databases, creates a new connection each time (thread-safe).
+ """
+ if self.db_path == ":memory:":
+ if self._persistent_conn is None:
+ self._persistent_conn = sqlite3.connect(":memory:")
+ self._persistent_conn.row_factory = sqlite3.Row
+ self._persistent_conn.execute("PRAGMA foreign_keys=ON")
+ return self._persistent_conn
+ else:
+ conn = sqlite3.connect(self.db_path)
+ conn.row_factory = sqlite3.Row
+ conn.execute("PRAGMA journal_mode=WAL")
+ conn.execute("PRAGMA foreign_keys=ON")
+ return conn
+
+ def initialize(self) -> None:
+ """Create tables if they don't exist."""
+ with self._get_conn() as conn:
+ conn.executescript("""
+ CREATE TABLE IF NOT EXISTS conversations (
+ id TEXT PRIMARY KEY,
+ title TEXT NOT NULL DEFAULT 'New Chat',
+ created_at REAL NOT NULL,
+ updated_at REAL NOT NULL,
+ metadata TEXT -- JSON-serialized dict
+ );
+
+ CREATE TABLE IF NOT EXISTS messages (
+ id TEXT PRIMARY KEY,
+ conversation_id TEXT NOT NULL,
+ role TEXT NOT NULL CHECK (role IN ('user', 'assistant')),
+ content TEXT NOT NULL,
+ timestamp REAL NOT NULL,
+ recall_topk TEXT, -- JSON-serialized dict
+ FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE
+ );
+
+ CREATE INDEX IF NOT EXISTS idx_messages_conversation
+ ON messages(conversation_id, timestamp);
+
+ CREATE INDEX IF NOT EXISTS idx_conversations_updated
+ ON conversations(updated_at DESC);
+ """)
+
+ def save_conversation(self, conv: Conversation) -> None:
+ """Insert or update conversation metadata."""
+ metadata_json = json.dumps(conv.metadata) if conv.metadata else None
+ with self._get_conn() as conn:
+ conn.execute(
+ """INSERT OR REPLACE INTO conversations (id, title, created_at, updated_at, metadata)
+ VALUES (?, ?, ?, ?, ?)""",
+ (conv.id, conv.title, conv.created_at, conv.updated_at, metadata_json),
+ )
+
+ def load_conversation(self, conv_id: str) -> Optional[Conversation]:
+ """Load a full conversation with all messages."""
+ with self._get_conn() as conn:
+ row = conn.execute(
+ "SELECT * FROM conversations WHERE id = ?", (conv_id,)
+ ).fetchone()
+ if row is None:
+ return None
+
+ messages = self._load_messages(conn, conv_id)
+ metadata = json.loads(row["metadata"]) if row["metadata"] else None
+
+ return Conversation(
+ id=row["id"],
+ title=row["title"],
+ created_at=row["created_at"],
+ updated_at=row["updated_at"],
+ messages=messages,
+ metadata=metadata,
+ )
+
+ def list_conversations(self) -> List[Conversation]:
+ """List all conversations (metadata only, no messages). Sorted by updated_at DESC."""
+ with self._get_conn() as conn:
+ rows = conn.execute(
+ "SELECT * FROM conversations ORDER BY updated_at DESC"
+ ).fetchall()
+
+ conversations = []
+ for row in rows:
+ metadata = json.loads(row["metadata"]) if row["metadata"] else None
+ # Count messages for the summary
+ msg_count = conn.execute(
+ "SELECT COUNT(*) as cnt FROM messages WHERE conversation_id = ?",
+ (row["id"],),
+ ).fetchone()["cnt"]
+ turn_count = conn.execute(
+ "SELECT COUNT(*) as cnt FROM messages WHERE conversation_id = ? AND role = 'user'",
+ (row["id"],),
+ ).fetchone()["cnt"]
+
+ conv = Conversation(
+ id=row["id"],
+ title=row["title"],
+ created_at=row["created_at"],
+ updated_at=row["updated_at"],
+ messages=[], # Lightweight: no messages loaded
+ metadata=metadata,
+ )
+ # Attach counts as metadata for API convenience
+ conv._message_count = msg_count # type: ignore[attr-defined]
+ conv._turn_count = turn_count # type: ignore[attr-defined]
+ conversations.append(conv)
+
+ return conversations
+
+ def delete_conversation(self, conv_id: str) -> bool:
+ """Delete a conversation and all its messages."""
+ with self._get_conn() as conn:
+ cursor = conn.execute("DELETE FROM conversations WHERE id = ?", (conv_id,))
+ return cursor.rowcount > 0
+
+ def add_message(self, conv_id: str, message: Message) -> None:
+ """Append a message to an existing conversation."""
+ recall_topk_json = json.dumps(message.recall_topk) if message.recall_topk else None
+ with self._get_conn() as conn:
+ conn.execute(
+ """INSERT INTO messages (id, conversation_id, role, content, timestamp, recall_topk)
+ VALUES (?, ?, ?, ?, ?, ?)""",
+ (message.id, conv_id, message.role, message.content, message.timestamp, recall_topk_json),
+ )
+ # Update conversation's updated_at
+ conn.execute(
+ "UPDATE conversations SET updated_at = ? WHERE id = ?",
+ (message.timestamp, conv_id),
+ )
+
+ def get_messages(self, conv_id: str) -> List[Message]:
+ """Get all messages for a conversation, ordered by timestamp."""
+ with self._get_conn() as conn:
+ return self._load_messages(conn, conv_id)
+
+ def _load_messages(self, conn: sqlite3.Connection, conv_id: str) -> List[Message]:
+ """Internal: load messages from an existing connection."""
+ rows = conn.execute(
+ "SELECT * FROM messages WHERE conversation_id = ? ORDER BY timestamp ASC",
+ (conv_id,),
+ ).fetchall()
+
+ messages = []
+ for row in rows:
+ recall_topk = json.loads(row["recall_topk"]) if row["recall_topk"] else None
+ messages.append(Message(
+ id=row["id"],
+ role=row["role"],
+ content=row["content"],
+ timestamp=row["timestamp"],
+ recall_topk=recall_topk,
+ ))
+ return messages
+
+ def close(self) -> None:
+ """Close the persistent connection if using :memory:."""
+ if self._persistent_conn is not None:
+ self._persistent_conn.close()
+ self._persistent_conn = None
diff --git a/scripts/MSA.sh b/scripts/MSA.sh
new file mode 100644
index 0000000..b04b730
--- /dev/null
+++ b/scripts/MSA.sh
@@ -0,0 +1,93 @@
+#!/bin/bash
+export MASTER_PORT=29509
+export CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7 # 在这里设置你需要使用的Device
+# export CUDA_VISIBLE_DEVICES=0,1
+
+# model
+model_path=ckpt/MSA-4B
+
+top_p=0.9
+temperature=0.0
+max_length=2048
+template=QWEN3_INSTRUCT_TEMPLATE
+
+
+# Statistics
+total_benchmarks=${#benchmarks[@]}
+current=0
+success_count=0
+fail_count=0
+failed_benchmarks=()
+
+echo "=========================================="
+echo "Start running all benchmark evaluations"
+echo "Total: ${total_benchmarks} benchmarks"
+echo "Log directory: $log_dir"
+echo "=========================================="
+echo ""
+
+# Run each benchmark
+for entry in "${benchmarks[@]}"; do
+ benchmark="${entry%%:*}"
+ batch_size="${entry##*:}"
+ current=$((current + 1))
+ echo "Start time: $(date '+%Y-%m-%d %H:%M:%S')"
+
+ # Create separate log file for each benchmark
+ log_file="$log_dir/${benchmark}.log"
+ json_file="$log_dir/${benchmark}.json"
+
+ # Run evaluation and record logs
+ python -u src/app/benchmark.py \
+ --benchmark "$benchmark" \
+ --model_path "$model_path" \
+ --top_p "$top_p" \
+ --temperature "$temperature" \
+ --max_length "$max_length" \
+ --template "$template" \
+ --output_file "$json_file" \
+ --max_batch_size "$batch_size" \
+ --max_chunk_per_block 16384 \
+ --block_size 2048 \
+ 2>&1 | tee $log_file
+
+ # Check exit status
+ exit_code=${PIPESTATUS[0]}
+
+ if [ $exit_code -eq 0 ]; then
+ echo "[$current/$total_benchmarks] $benchmark finished (success)"
+ # Print benchmark name and metrics
+ echo "========== $benchmark Results =========="
+ python -c "import json; d=json.load(open('$json_file')); [print(f' {k}: {v}') for k,v in d.get(list(d.keys())[0],{}).get('precision',{}).get('metrics',{}).items()]" 2>/dev/null || echo " (failed to parse metrics)"
+ echo "========================================"
+ success_count=$((success_count + 1))
+ else
+ echo "[$current/$total_benchmarks] $benchmark failed (exit code: $exit_code)"
+ fail_count=$((fail_count + 1))
+ failed_benchmarks+=("$benchmark")
+ fi
+
+ echo "End time: $(date '+%Y-%m-%d %H:%M:%S')"
+ echo "----------------------------------------"
+ echo ""
+done
+
+# Summary
+echo "=========================================="
+echo "All benchmark evaluations completed"
+echo "=========================================="
+echo "Total: $total_benchmarks"
+echo "Success: $success_count"
+echo "Failed: $fail_count"
+echo ""
+
+if [ $fail_count -gt 0 ]; then
+ echo "Failed benchmarks:"
+ for failed in "${failed_benchmarks[@]}"; do
+ echo " - $failed"
+ done
+ echo ""
+fi
+
+echo "All logs saved in: $log_dir"
+echo "=========================================="
\ No newline at end of file
diff --git a/scripts/simple_msa_demo.py b/scripts/simple_msa_demo.py
new file mode 100644
index 0000000..1f721ae
--- /dev/null
+++ b/scripts/simple_msa_demo.py
@@ -0,0 +1,777 @@
+#!/usr/bin/env python3
+"""
+simple_msa_demo.py — MSA 完整推理流程示例
+
+功能:
+ 1. 读取指定目录下的 .txt 文档
+ 2. Stage 1: 用 MSA 模型编码文档为 chunk-pooled KV cache
+ 3. Stage 2: 用 router_q_proj/router_k_proj 做完整路由检索
+ 4. Stage 3: 拼接 template prefix + 选中文档 KV + query,调用 model.generate()
+
+用法:
+ python scripts/simple_msa_demo.py \
+ --model_path ckpt/MSA-4B \
+ --doc_dir /path/to/your/documents \
+ --query "你的问题"
+
+与项目完整流程的对应关系:
+ 本脚本 = PrefillStage1Worker._inference() (Stage 1)
+ + Memory.prefill_stage2() (Stage 2 路由打分)
+ + MSAService.generate() (Stage 3 生成)
+ 去掉了多 GPU 通信 (NCCL all-gather/all-to-all),仅用单 GPU 运行。
+"""
+
+import argparse
+import os
+import sys
+import pathlib
+import glob
+from typing import List, Dict, Tuple, Optional
+
+import torch
+import torch.nn.functional as F
+import numpy as np
+from transformers import AutoTokenizer
+
+project_path = pathlib.Path(__file__).parent.parent
+sys.path.insert(0, str(project_path))
+
+from src.msa.model import MSAForCausalLM, MSAConfig
+from src.msa.memory_sparse_attention import MemorySparseAttention
+from src.utils.cache import CustomDynamicCache, create_cache
+from src.utils.template import QWEN3_INSTRUCT_TEMPLATE
+
+
+# ============================================================
+# 1. 读取文档
+# ============================================================
+
+def _read_text_file(path: str) -> str:
+ """尝试多种编码读取文本文件。"""
+ for encoding in ["utf-8", "gbk", "gb2312", "gb18030", "latin-1"]:
+ try:
+ with open(path, "r", encoding=encoding) as f:
+ return f.read()
+ except (UnicodeDecodeError, UnicodeError):
+ continue
+ raise UnicodeDecodeError(f"无法识别文件编码: {path}")
+
+
+def load_documents(doc_dir: str) -> List[str]:
+ """读取目录下所有 .txt 文件,返回文档内容列表。"""
+ docs = []
+ paths = sorted(glob.glob(os.path.join(doc_dir, "*.txt")))
+ if not paths:
+ raise FileNotFoundError(f"在 {doc_dir} 下未找到 .txt 文件")
+ for path in paths:
+ content = _read_text_file(path).strip()
+ if content:
+ docs.append(content)
+ print(f"[Stage 0] 加载了 {len(docs)} 篇文档,来自 {doc_dir}")
+ return docs
+
+
+# ============================================================
+# 2. Stage 1: 编码文档为 KV cache(完整实现)
+# ============================================================
+
+def encode_documents(
+ model: MSAForCausalLM,
+ tokenizer,
+ documents: List[str],
+ pooling_kernel_size: int,
+ device: str,
+) -> Dict:
+ """
+ Stage 1 完整实现:对所有文档做 forward,得到 chunk-pooled KV cache。
+
+ 对应原始代码:
+ PrefillStage1Worker._inference() → MemorySparseAttention.forward(stage="prefill_stage1")
+
+ 返回一个 dict 包含所有推理产物:
+ - kv_caches: {layer_idx: (K, V)} 每层的 chunk-pooled KV
+ - router_k_caches: {layer_idx: rk} 每层的 router key (decouple_router=True 时)
+ - template_prefix_kvcache: {layer_idx: (k, v)} 模板前缀 KV
+ - pooled_doc_ids: {layer_idx: tensor} 每个 chunk 对应的文档 ID
+ - chunk_sizes: 每篇文档的 chunk 数量
+ """
+ msa_config = model.config.msa_config
+ router_layer_idx = msa_config.router_layer_idx
+ if router_layer_idx == "all":
+ router_layers = list(range(model.config.num_hidden_layers))
+ else:
+ router_layers = [int(i) for i in router_layer_idx.split(",")]
+
+ # ── 拼接所有文档为一个序列 ──
+ # 对应 PrefillStage1Worker._prepare_block_inputs()
+ all_input_ids = []
+ all_attention_mask = []
+ all_doc_ids = []
+ all_position_ids = []
+ chunk_sizes = []
+
+ for doc_idx, doc_text in enumerate(documents):
+ doc_inputs = tokenizer(doc_text, add_special_tokens=False)
+ doc_token_ids = doc_inputs["input_ids"]
+ length = len(doc_token_ids)
+
+ all_input_ids.extend(doc_token_ids)
+ all_attention_mask.extend([1] * length)
+ # doc_id 从 1 开始(0 留给 query 区域)
+ all_doc_ids.extend([doc_idx + 1] * length)
+ all_position_ids.extend(list(range(length)))
+
+ n_chunks = (length + pooling_kernel_size - 1) // pooling_kernel_size
+ chunk_sizes.append(n_chunks)
+
+ input_ids = torch.LongTensor([all_input_ids]).to(device)
+ attention_mask = torch.LongTensor([all_attention_mask]).to(device)
+ doc_ids = torch.LongTensor([all_doc_ids]).to(device)
+ position_ids = torch.LongTensor([all_position_ids]).to(device)
+
+ # ── 创建 Cache ──
+ past_key_values = CustomDynamicCache()
+ n_layers = model.config.num_hidden_layers
+ for layer_idx in range(n_layers):
+ past_key_values.record_kwargs(layer_idx, {"stage": "prefill_stage1"})
+
+ print(f"[Stage 1] 正在编码文档... (共 {sum(chunk_sizes)} chunks, {len(documents)} 篇文档)")
+
+ with torch.no_grad():
+ outputs = model.model(
+ input_ids=input_ids,
+ attention_mask=attention_mask,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ use_cache=True,
+ doc_ids=doc_ids,
+ )
+
+ # ── 提取产物 ──
+ past_kv = outputs.past_key_values
+ result = {
+ "kv_caches": {},
+ "router_k_caches": {},
+ "template_prefix_kvcache": {},
+ "pooled_doc_ids": {},
+ "chunk_sizes": chunk_sizes,
+ }
+
+ for layer_idx in range(n_layers):
+ k, v = past_kv.get_kvcache(layer_idx)
+ result["kv_caches"][layer_idx] = (k, v)
+
+ # router key (decouple_router=True 时独立存在)
+ rk = past_kv.get_router_kcache(layer_idx)
+ if rk is not None:
+ result["router_k_caches"][layer_idx] = rk
+
+ # 模板前缀 KV
+ kwargs = past_kv.cache_kwargs.get(layer_idx, {})
+ if "template_prefix_kcache" in kwargs:
+ result["template_prefix_kvcache"][layer_idx] = (
+ kwargs["template_prefix_kcache"].to(device),
+ kwargs["template_prefix_vcache"].to(device),
+ )
+
+ # pooled doc ids(每个 chunk 属于哪篇文档)
+ if "pooled_doc_ids" in kwargs:
+ result["pooled_doc_ids"][layer_idx] = kwargs["pooled_doc_ids"]
+
+ print(f"[Stage 1] 编码完成,KV cache 形状: {result['kv_caches'][router_layers[0]][0].shape}")
+ return result
+
+
+# ============================================================
+# 3. Stage 2: 完整路由检索
+# ============================================================
+
+def retrieve_documents(
+ model: MSAForCausalLM,
+ tokenizer,
+ query: str,
+ stage1_result: Dict,
+ top_k: int,
+ device: str,
+) -> Dict:
+ """
+ Stage 2 完整实现:用 router_q_proj/router_k_proj 做路由检索。
+
+ 对应原始代码:
+ MemorySparseAttention.forward(stage="prefill_stage2")
+ → Memory.prefill_stage2() (路由打分)
+ → Memory.doc_query() (选中文档 KV 提取)
+
+ 完整流程:
+ 1. 对 query 做 forward,经过 router layer 时:
+ a. router_q_proj 投影 query → routing_q
+ b. routing_q 与 router_k (或 pooled K) 做 cosine similarity
+ c. head_reduce → query_reduce → chunk_reduce 得到文档分数
+ d. top-k 选出文档
+ 2. 提取选中文档的 K/V
+
+ 返回:
+ - selected_k: [1, n_kv_heads, selected_chunks, head_dim]
+ - selected_v: [1, n_kv_heads, selected_chunks, head_dim]
+ - scores: [1, top_k]
+ - selected_doc_ids: 选中的文档 ID 列表
+ - template_prefix_kvcache: 模板前缀 KV
+ """
+ msa_config = model.config.msa_config
+ router_layer_idx = msa_config.router_layer_idx
+ if router_layer_idx == "all":
+ router_layers = list(range(model.config.num_hidden_layers))
+ else:
+ router_layers = [int(i) for i in router_layer_idx.split(",")]
+
+ # 用第一个 router layer 做检索
+ layer_idx = router_layers[0]
+
+ # 获取该层的 attention module
+ attn_layer = model.model.layers[layer_idx].self_attn
+ assert isinstance(attn_layer, MemorySparseAttention), \
+ f"Layer {layer_idx} is not MemorySparseAttention, got {type(attn_layer)}"
+
+ # ── 构建 query 输入 ──
+ # 对应 MSAService._apply_template()
+ template = QWEN3_INSTRUCT_TEMPLATE
+ prompt_text = (
+ "\nPlease answer the question based on the above historical document information\n\n"
+ + query
+ + "\nPlease return all documents related to the question\n"
+ )
+
+ prompt_inputs = tokenizer(prompt_text, add_special_tokens=False)
+ prompt_ids = prompt_inputs["input_ids"]
+
+ # 模板后缀
+ pad_token = tokenizer.pad_token
+ pad_token_id = tokenizer.pad_token_id
+ template_str = template["prompt"].replace("{prompt}", pad_token)
+ template_inputs = tokenizer(template_str, add_special_tokens=False)
+ pad_index = template_inputs["input_ids"].index(pad_token_id)
+ tail_ids = template_inputs["input_ids"][pad_index + 1:]
+
+ # response head
+ response_head = "<|im_start|>"
+ response_head_ids = tokenizer(response_head, add_special_tokens=False)["input_ids"]
+
+ # 拼接完整序列
+ full_ids = prompt_ids + tail_ids + response_head_ids
+ input_ids_tensor = torch.LongTensor([full_ids]).to(device)
+ seq_len = input_ids_tensor.shape[1]
+
+ attention_mask = torch.ones(1, seq_len, dtype=torch.long, device=device)
+
+ # doc_ids: query 区域 = 0, template 后缀 = -2, response head = -1
+ doc_ids = (
+ [0] * len(prompt_ids)
+ + [-2] * len(tail_ids)
+ + [-1] * len(response_head_ids)
+ )
+ doc_ids_tensor = torch.LongTensor([doc_ids]).to(device)
+
+ # position ids(Global RoPE:从 top_k + 3 开始)
+ top_k_actual = min(top_k, stage1_result["kv_caches"][layer_idx][0].shape[2])
+ start_pos = 3 + top_k_actual
+ position_ids = list(range(start_pos, start_pos + seq_len))
+ position_ids_tensor = torch.LongTensor([position_ids]).to(device)
+
+ # ── 创建 Cache 并注入 Stage 1 产物 ──
+ past_key_values = CustomDynamicCache()
+ n_layers = model.config.num_hidden_layers
+ for li in range(n_layers):
+ kwargs = {"stage": "prefill_stage2"}
+ # 注入 template prefix KV
+ if li in stage1_result["template_prefix_kvcache"]:
+ t_k, t_v = stage1_result["template_prefix_kvcache"][li]
+ kwargs["template_prefix_kcache"] = t_k
+ kwargs["template_prefix_vcache"] = t_v
+ past_key_values.record_kwargs(li, kwargs)
+
+ # 将 Stage 1 的 KV cache 注入到 past_key_values
+ for li in range(n_layers):
+ if li in stage1_result["kv_caches"]:
+ k, v = stage1_result["kv_caches"][li]
+ past_key_values.update(k, v, li)
+ if li in stage1_result["router_k_caches"]:
+ rk = stage1_result["router_k_caches"][li]
+ past_key_values.update_router_kcache(rk, li)
+ # 注入 pooled_doc_ids
+ if li in stage1_result["pooled_doc_ids"]:
+ past_key_values.cache_kwargs[li]["pooled_doc_ids"] = stage1_result["pooled_doc_ids"][li]
+
+ # ── 执行 forward(会触发 MemorySparseAttention 的 prefill_stage2 分支)──
+ print(f"[Stage 2] 正在路由检索... (query length: {seq_len}, top_k: {top_k_actual})")
+
+ # 设置 memory_client 让 MemorySparseAttention 能访问 BlockData
+ # 对应 MSAService.setup_memory_client()
+ _setup_memory_client_for_demo(
+ model=model,
+ kv_caches=stage1_result["kv_caches"],
+ router_k_caches=stage1_result["router_k_caches"],
+ template_prefix_kvcache=stage1_result["template_prefix_kvcache"],
+ pooled_doc_ids=stage1_result["pooled_doc_ids"],
+ chunk_sizes=stage1_result["chunk_sizes"],
+ device=device,
+ )
+
+ with torch.no_grad():
+ outputs = model.model(
+ input_ids=input_ids_tensor,
+ attention_mask=attention_mask,
+ position_ids=position_ids_tensor,
+ past_key_values=past_key_values,
+ use_cache=True,
+ doc_ids=doc_ids_tensor,
+ )
+
+ # ── 提取路由结果 ──
+ # 从 past_key_values 的 cache_kwargs 中提取 recall_topk
+ # 对应 MemorySparseAttention.forward(stage="prefill_stage2") 中的 recall_topk 逻辑
+ recall_topk = past_key_values.cache_kwargs[layer_idx].get("recall_topk", None)
+
+ # 提取 compacked KV cache(已组装好的 template + selected docs + query)
+ compacked_k = past_key_values.cache_kwargs[layer_idx].get("compacked_key_cache", None)
+ compacked_v = past_key_values.cache_kwargs[layer_idx].get("compacked_value_cache", None)
+
+ result = {
+ "past_key_values": past_key_values,
+ "compacked_k": compacked_k,
+ "compacked_v": compacked_v,
+ "recall_topk": recall_topk,
+ "input_ids": input_ids_tensor,
+ "attention_mask": attention_mask,
+ "doc_ids": doc_ids_tensor,
+ "position_ids": position_ids_tensor,
+ }
+
+ if recall_topk:
+ for item in recall_topk:
+ doc_ids_list = item.get("topk_doc_ids", [])
+ scores_list = item.get("score", [])
+ print(f"[Stage 2] 选出 {len(doc_ids_list)} 个文档 chunks, 最高分: {max(scores_list):.4f}")
+ else:
+ print("[Stage 2] 路由完成(无 recall_topk 输出)")
+
+ return result
+
+
+def _setup_memory_client_for_demo(
+ model: MSAForCausalLM,
+ kv_caches: Dict,
+ router_k_caches: Dict,
+ template_prefix_kvcache: Dict,
+ pooled_doc_ids: Dict,
+ chunk_sizes: List[int],
+ device: str,
+):
+ """
+ 为 MemorySparseAttention 设置 memory_client,使其能访问文档 KV cache。
+
+ 对应 MSAService.setup_memory_client()。
+ 在完整流程中,MemorySparseAttention 通过 self.memory_client.doc_query()
+ 访问 BlockData 中的 K/V/rk。这里我们构造一个最小实现。
+ """
+ # 计算文档元数据
+ nr_docs = len(chunk_sizes)
+ doc_lens = [0] + chunk_sizes # 第 0 位是 padding
+ doc_offsets = [0]
+ for cl in chunk_sizes:
+ doc_offsets.append(doc_offsets[-1] + cl)
+ doc_ids_list = list(range(nr_docs + 1)) # [0, 1, 2, ..., nr_docs]
+
+ doc_lens_cpu = torch.LongTensor(doc_lens)
+ doc_offsets_cpu = torch.LongTensor(doc_offsets)
+ doc_ids_tensor = torch.LongTensor(doc_ids_list).to(device)
+
+ # 构造 k_slices 和 slice_desc(对应 Memory._build_k_slices)
+ n_layers = len(kv_caches)
+ k_slices = {}
+ slice_desc_list = []
+
+ for layer_idx in range(n_layers):
+ if layer_idx not in kv_caches:
+ continue
+ k, v = kv_caches[layer_idx]
+ # k shape: [1, n_kv_heads, n_chunks, head_dim]
+ n_chunks = k.shape[2]
+ # 构造 k_slice_t: [1, n_kv_heads, 1, head_dim, n_chunks]
+ k_slice_t = k.permute(0, 1, 3, 2).unsqueeze(2)
+ k_slices[layer_idx] = [k_slice_t]
+
+ # 简化的 slice_desc
+ class SimpleSliceDesc:
+ def __init__(self, nr_chunks, nr_docs, doc_ids):
+ self.nr_chunks = nr_chunks
+ self.nr_docs = nr_docs
+ self.local_doc_ids_0 = torch.arange(nr_docs + 1, device=device).unsqueeze(0)
+ self.original_doc_ids = doc_ids.unsqueeze(0)
+ self.global_doc_ids = doc_ids.unsqueeze(0)
+ slice_desc_list.append(SimpleSliceDesc(n_chunks, nr_docs, doc_ids_tensor))
+
+ class DemoMemoryClient:
+ """最小 memory_client 实现,让 MemorySparseAttention 能访问文档 KV。"""
+ def __init__(self):
+ self.blocks = {}
+ self.block_desc = type('BlockDesc', (), {
+ 'nr_docs': nr_docs,
+ 'doc_lens_cpu': doc_lens_cpu,
+ 'doc_offsets_cpu': doc_offsets_cpu,
+ 'doc_ids': doc_ids_tensor,
+ 'doc_ids_cpu': torch.LongTensor(doc_ids_list),
+ })()
+ self.k_slices = k_slices
+ self.slice_desc = slice_desc_list
+ self.model_config = type('MC', (), {'doc_top_k': 16})()
+ self.template_prefix_kvcache = template_prefix_kvcache
+ self.pooled_doc_ids = pooled_doc_ids
+ self.device = device
+ self.world_size = 1
+ self.num_key_value_groups = (
+ model.config.num_attention_heads // model.config.num_key_value_heads
+ )
+ self.head_reduce_method = msa_config.head_reduce_method
+ self.query_reduce_method = msa_config.query_reduce_method
+ self.chunk_reduce_method = msa_config.chunk_reduce_method
+ self.decouple_router = msa_config.decouple_router
+ self.scaling = -1.0 if "INFONCE" in msa_config.aux_loss_method and self.decouple_router else 1.0
+
+ # 初始化 blocks
+ for li in range(n_layers):
+ from src.msa_service import BlockData
+ block_data = BlockData()
+ if li in kv_caches:
+ k, v = kv_caches[li]
+ block_data.k = k
+ block_data.v = v
+ if li in router_k_caches:
+ block_data.rk = router_k_caches[li]
+ self.blocks[li] = block_data
+
+ # 构建 k_slices 和 slice_desc(对应 Memory._build_k_slices)
+ self._build_k_slices()
+
+ def _build_k_slices(self):
+ """构建分片 K slices 用于路由打分。"""
+ msa_config_local = model.config.msa_config
+ router_layer_idx_str = msa_config_local.router_layer_idx
+ if router_layer_idx_str == "all":
+ router_layers_local = list(range(model.config.num_hidden_layers))
+ else:
+ router_layers_local = [int(i) for i in router_layer_idx_str.split(",")]
+
+ self.k_slices = {}
+ self.slice_desc = []
+ slice_chunk_size = 16 * 1024 # 默认值
+
+ for li in router_layers_local:
+ if li not in self.blocks:
+ continue
+ block = self.blocks[li]
+ router_k = block.get_router_k() # rk if exists, else k
+ if router_k is None:
+ continue
+
+ n_kv_heads = router_k.shape[1]
+ n_chunks = router_k.shape[2]
+
+ # 按 slice_chunk_size 分片
+ chunks_done = 0
+ k_slices_for_layer = []
+ slice_descs = []
+ while chunks_done < n_chunks:
+ this_chunk = min(slice_chunk_size, n_chunks - chunks_done)
+ k_slice = router_k[:, :, chunks_done:chunks_done + this_chunk, :]
+ # k_slice_t: [1, n_kv_heads, 1, head_dim, this_chunk]
+ k_slice_t = k_slice.permute(0, 1, 3, 2).unsqueeze(2)
+ k_slices_for_layer.append(k_slice_t)
+
+ # 构建 slice 描述
+ doc_ids_for_slice = self.pooled_doc_ids.get(li)
+ if doc_ids_for_slice is not None:
+ slice_doc_ids = doc_ids_for_slice[chunks_done:chunks_done + this_chunk]
+ else:
+ slice_doc_ids = torch.arange(1, this_chunk + 1, device=self.device)
+
+ class SliceDesc:
+ pass
+ sd = SliceDesc()
+ sd.nr_chunks = this_chunk
+ sd.nr_docs = self.block_desc.nr_docs
+
+ # local_doc_ids_0: [1, nr_chunks] — 每个 chunk 对应的 local doc id
+ sd.local_doc_ids_0 = slice_doc_ids.unsqueeze(0)
+ # original_doc_ids: [1, nr_docs] — local → global 映射
+ sd.original_doc_ids = self.block_desc.doc_ids.unsqueeze(0)
+ slice_descs.append(sd)
+
+ chunks_done += this_chunk
+
+ self.k_slices[li] = k_slices_for_layer
+ if not self.slice_desc:
+ self.slice_desc = slice_descs
+
+ def get_template_prefix_kvcaches(self, layer_idx):
+ return self.template_prefix_kvcache[layer_idx]
+
+ def doc_query(self, query_states, query_mask, layer_idx):
+ """
+ 单 GPU 版 doc_query。
+ 对应 MSAService.doc_query(),去掉了 NCCL 通信。
+ """
+ bsz, nhead, seqlen, hdim = query_states.shape
+ num_kv_groups = self.num_key_value_groups
+ top_k = self.model_config.doc_top_k
+ dtype = query_states.dtype
+
+ # ── Phase 1: 本地打分 ──
+ # 对应 Memory.prefill_stage2()
+ total_docs = self.block_desc.nr_docs + 1
+ min_val = torch.finfo(dtype).min
+ global_doc_scores = torch.full((bsz, total_docs), min_val, dtype=dtype, device=self.device)
+
+ # reshape query for matmul
+ query_reshaped = query_states.view(bsz, self.num_key_value_groups, -1, seqlen, hdim) * self.scaling
+
+ for slice_desc, k_slice_t in zip(self.slice_desc, self.k_slices[layer_idx]):
+ # attn_scores: [bsz, n_kv_heads, n_groups, seqlen, chunk]
+ attn_scores = torch.matmul(query_reshaped, k_slice_t)
+
+ # mask invalid query positions
+ routing_mask = ~query_mask
+ attn_scores.masked_fill_(routing_mask, min_val)
+
+ # head_reduce
+ if self.head_reduce_method == "max":
+ scores = attn_scores.flatten(1, 2).max(dim=1).values # [bsz, seqlen, chunk]
+ elif self.head_reduce_method == "mean":
+ scores = attn_scores.flatten(1, 2).mean(dim=1)
+ else:
+ scores = attn_scores.flatten(1, 2).max(dim=1).values
+
+ # query_reduce
+ if self.query_reduce_method == "max":
+ scores = scores.max(dim=1).values # [bsz, chunk]
+ elif self.query_reduce_method == "mean":
+ valid_mask = query_mask.squeeze(2).squeeze(1) # [bsz, seqlen]
+ scores_clean = torch.where(valid_mask.unsqueeze(-1), scores, torch.zeros_like(scores))
+ counts = valid_mask.sum(dim=1, keepdim=True).clamp(min=1.0)
+ scores = scores_clean.sum(dim=1) / counts
+ else:
+ scores = scores.max(dim=1).values
+
+ # chunk → doc scatter
+ scatter_indices = slice_desc.local_doc_ids_0.expand(bsz, -1)
+ local_doc_scores = torch.full((bsz, slice_desc.nr_docs + 1), min_val, device=self.device, dtype=dtype)
+
+ if self.chunk_reduce_method == "max":
+ local_doc_scores.scatter_reduce_(dim=1, index=scatter_indices, src=scores, reduce="amax", include_self=True)
+ elif self.chunk_reduce_method == "mean":
+ doc_sums = torch.zeros_like(local_doc_scores)
+ doc_sums.scatter_reduce_(dim=1, index=scatter_indices, src=scores, reduce="sum", include_self=False)
+ doc_counts = torch.zeros_like(local_doc_scores)
+ ones = torch.ones_like(scores)
+ doc_counts.scatter_reduce_(dim=1, index=scatter_indices, src=ones, reduce="sum", include_self=False)
+ doc_counts = doc_counts.clamp(min=1.0)
+ mean_scores = doc_sums / doc_counts
+ local_doc_scores = torch.where(doc_counts > 0, mean_scores, local_doc_scores)
+ else:
+ local_doc_scores.scatter_reduce_(dim=1, index=scatter_indices, src=scores, reduce="amax", include_self=True)
+
+ # local → global
+ global_indices = slice_desc.original_doc_ids.expand(bsz, -1)
+ global_doc_scores.scatter_reduce_(dim=1, index=global_indices, src=local_doc_scores, reduce="amax", include_self=True)
+
+ # ── Phase 2: Top-K 选择 ──
+ k = min(top_k, global_doc_scores.shape[1])
+ final_scores, batch_selected_local_ids = torch.topk(global_doc_scores, k=k, dim=1)
+ batch_selected_global_ids = self.block_desc.doc_ids[batch_selected_local_ids]
+
+ # ── Phase 3: 提取选中文档的 K/V ──
+ block = self.blocks[layer_idx]
+ doc_k = block.k # [1, n_kv_heads, n_chunks, head_dim]
+ doc_v = block.v
+
+ # 收集选中 chunks 的 K/V
+ # 需要从 pooled_doc_ids 中找到被选中文档对应的所有 chunks
+ pooled_ids = self.pooled_doc_ids.get(layer_idx)
+ if pooled_ids is not None:
+ # 找到属于选中文档的所有 chunk 索引
+ selected_doc_set = set(batch_selected_global_ids[0].cpu().tolist())
+ chunk_mask = torch.tensor(
+ [pooled_ids[i].item() in selected_doc_set for i in range(pooled_ids.shape[0])],
+ device=self.device,
+ )
+ selected_chunk_indices = chunk_mask.nonzero(as_tuple=False).squeeze(-1)
+
+ final_k = doc_k[:, :, selected_chunk_indices, :]
+ final_v = doc_v[:, :, selected_chunk_indices, :]
+ final_selected_doc_ids = batch_selected_global_ids
+ else:
+ final_k = doc_k
+ final_v = doc_v
+ final_selected_doc_ids = batch_selected_global_ids
+
+ num_selected = torch.tensor([final_k.shape[2]], device=self.device, dtype=torch.long)
+
+ return final_k, final_v, final_scores, num_selected, final_selected_doc_ids
+
+ # 注入 memory_client 到所有 router layers
+ msa_config = model.config.msa_config
+ client = DemoMemoryClient()
+ for layer in model.model.modules():
+ if isinstance(layer, MemorySparseAttention):
+ layer.set_memory_client(client)
+
+
+# ============================================================
+# 4. Stage 3: 用 model.generate() 生成回答
+# ============================================================
+
+def generate_answer(
+ model: MSAForCausalLM,
+ tokenizer,
+ stage2_result: Dict,
+ max_new_tokens: int = 256,
+ temperature: float = 0.0,
+ top_p: float = 0.9,
+) -> str:
+ """
+ Stage 3: 调用 model.generate() 生成回答。
+
+ Stage 2 的 MemorySparseAttention 已经将 template prefix KV + 选中文档 KV
+ + query KV 组装成了 compacked_key_cache / compacked_value_cache,
+ 存在 past_key_values 的 cache_kwargs 中。
+
+ Stage 3 的 generate 过程中,MemorySparseAttention 的 "generate" 分支
+ 会直接使用这个 compacked cache 做 attention。
+ """
+ past_key_values = stage2_result["past_key_values"]
+ input_ids = stage2_result["input_ids"]
+ attention_mask = stage2_result["attention_mask"]
+ doc_ids = stage2_result["doc_ids"]
+ position_ids = stage2_result["position_ids"]
+
+ # 将 stage 切换为 "generate"
+ n_layers = model.config.num_hidden_layers
+ for layer_idx in range(n_layers):
+ past_key_values.record_kwargs(layer_idx, {"stage": "generate"})
+
+ # 重新注入 compacked cache 和其他 kwargs
+ # (record_kwargs 会覆盖旧值,需要保留已有的 compacked cache)
+ for layer_idx in range(n_layers):
+ if layer_idx in stage2_result["past_key_values"].cache_kwargs:
+ old_kwargs = stage2_result["past_key_values"].cache_kwargs[layer_idx]
+ if "compacked_key_cache" in old_kwargs:
+ past_key_values.cache_kwargs[layer_idx]["compacked_key_cache"] = old_kwargs["compacked_key_cache"]
+ past_key_values.cache_kwargs[layer_idx]["compacked_value_cache"] = old_kwargs["compacked_value_cache"]
+ past_key_values.cache_kwargs[layer_idx]["kv_lengths"] = old_kwargs["kv_lengths"]
+ past_key_values.cache_kwargs[layer_idx]["attention_mask"] = old_kwargs["attention_mask"]
+
+ # 设置 generate 所需的 meta
+ past_key_values.meta["require_recall_topk"] = False
+ past_key_values.meta["qa_mode"] = True
+ past_key_values.meta["max_generate_tokens"] = max_new_tokens
+ past_key_values.meta["tokenizer"] = tokenizer
+ past_key_values.meta["idx_to_doc"] = {} # 单文档模式无需 doc reference
+ past_key_values.meta["pattern"] = r"\[(\d+)\]"
+ past_key_values.meta["response_string"] = [""]
+
+ print(f"[Stage 3] 正在生成回答... (max_new_tokens: {max_new_tokens})")
+
+ with torch.no_grad():
+ generated_ids = model.generate(
+ input_ids=input_ids,
+ attention_mask=attention_mask,
+ doc_ids=doc_ids,
+ position_ids=position_ids,
+ past_key_values=past_key_values,
+ max_new_tokens=max_new_tokens,
+ do_sample=(temperature > 0),
+ temperature=temperature if temperature > 0 else None,
+ top_p=top_p if temperature > 0 else None,
+ pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
+ )
+
+ response = tokenizer.decode(generated_ids[0], skip_special_tokens=False)
+ return response
+
+
+# ============================================================
+# 5. 主流程
+# ============================================================
+
+def main():
+ parser = argparse.ArgumentParser(description="MSA 完整推理流程演示")
+ parser.add_argument("--model_path", type=str, default="ckpt/MSA-4B", help="模型路径")
+ parser.add_argument("--doc_dir", type=str, required=True, help="文档目录(包含 .txt 文件)")
+ parser.add_argument("--query", type=str, required=True, help="用户问题")
+ parser.add_argument("--pooling_kernel_size", type=int, default=64, help="chunk pooling 窗口大小")
+ parser.add_argument("--top_k", type=int, default=16, help="检索的 top-k 文档数")
+ parser.add_argument("--max_new_tokens", type=int, default=256, help="最大生成 token 数")
+ parser.add_argument("--temperature", type=float, default=0.0, help="采样温度 (0=greedy)")
+ parser.add_argument("--top_p", type=float, default=0.9, help="核采样 top-p")
+ parser.add_argument("--device", type=str, default="cuda:0", help="设备")
+ args = parser.parse_args()
+
+ device = args.device
+ if not torch.cuda.is_available():
+ device = "cpu"
+ print("[WARNING] CUDA 不可用,使用 CPU(会很慢)")
+
+ # ── 加载模型 ──
+ print(f"[Init] 加载模型: {args.model_path}")
+ tokenizer = AutoTokenizer.from_pretrained(args.model_path)
+ model = MSAForCausalLM.from_pretrained(
+ args.model_path,
+ use_cache=True,
+ attn_implementation="flash_attention_2" if torch.cuda.is_available() else "eager",
+ torch_dtype= torch.bfloat16,
+ device_map=device,
+ )
+ model.eval()
+ print("[Init] 模型加载完成")
+
+ # ── 读取文档 ──
+ documents = load_documents(args.doc_dir)
+
+ # ── Stage 1: 编码文档 ──
+ stage1_result = encode_documents(
+ model=model,
+ tokenizer=tokenizer,
+ documents=documents,
+ pooling_kernel_size=args.pooling_kernel_size,
+ device=device,
+ )
+
+ # ── Stage 2: 路由检索 ──
+ stage2_result = retrieve_documents(
+ model=model,
+ tokenizer=tokenizer,
+ query=args.query,
+ stage1_result=stage1_result,
+ top_k=args.top_k,
+ device=device,
+ )
+
+ # ── Stage 3: 生成回答 ──
+ response = generate_answer(
+ model=model,
+ tokenizer=tokenizer,
+ stage2_result=stage2_result,
+ max_new_tokens=args.max_new_tokens,
+ temperature=args.temperature,
+ top_p=args.top_p,
+ )
+
+ # ── 输出 ──
+ print("\n" + "=" * 60)
+ print(f"问题: {args.query}")
+ print("=" * 60)
+ print(f"回答:\n{response}")
+ print("=" * 60)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/src/app/benchmark.py b/src/app/benchmark.py
index bdbbfbd..a1fd2b9 100644
--- a/src/app/benchmark.py
+++ b/src/app/benchmark.py
@@ -246,11 +246,11 @@ def msa_benchmark(args, data):
)
final_result = {}
- if args.output_file:
+ if args.output_file: # 看是否已经有答案了
try:
with open(args.output_file, 'r') as f:
- exist_result = json.load(f)
- final_result = exist_result[args.case_name]
+ exist_result = json.load(f) # 得到的答案
+ final_result = exist_result[args.case_name] # 最终答案是加载的已有的json文件
except:
pass