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