merge: pull origin/main and resolve batch size conflicts
This commit is contained in:
commit
2b80e8c41f
@ -741,11 +741,11 @@ async def update_thread(
|
||||
)
|
||||
|
||||
|
||||
|
||||
# ================================
|
||||
# > === 附件管理分组 ===
|
||||
# ================================
|
||||
|
||||
|
||||
@chat.post("/thread/{thread_id}/attachments", response_model=AttachmentResponse)
|
||||
async def upload_thread_attachment(
|
||||
thread_id: str,
|
||||
|
||||
@ -14,6 +14,7 @@ from src.agents.common.backends import create_agent_composite_backend
|
||||
from src.agents.common.middlewares import RuntimeConfigMiddleware, SummaryOffloadMiddleware, save_attachments_to_fs
|
||||
from src.agents.common.tools import get_tavily_search
|
||||
from src.services.mcp_service import get_tools_from_all_servers
|
||||
from src.utils import logger
|
||||
|
||||
|
||||
from .context import DeepContext
|
||||
@ -83,11 +84,14 @@ class DeepAgent(BaseAgent):
|
||||
if tavily_search:
|
||||
tools.append(tavily_search)
|
||||
|
||||
# Assert that search tool is available for DeepAgent
|
||||
assert tools, (
|
||||
"DeepAgent requires at least one search tool. "
|
||||
"Please configure TAVILY_API_KEY environment variable to enable web search."
|
||||
)
|
||||
# # Assert that search tool is available for DeepAgent
|
||||
# assert tools, (
|
||||
# "DeepAgent requires at least one search tool. "
|
||||
# "Please configure TAVILY_API_KEY environment variable to enable web search."
|
||||
# )
|
||||
if not tools:
|
||||
logger.warning("No search tools configured, DeepAgent will work without web search")
|
||||
tools = []
|
||||
return tools
|
||||
|
||||
async def get_graph(self, **kwargs):
|
||||
|
||||
@ -88,17 +88,32 @@ class BaseEmbeddingModel(ABC):
|
||||
task_id = hashstr(messages)
|
||||
self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0}
|
||||
|
||||
tasks = []
|
||||
# 保留原有逻辑:
|
||||
# 使用 asyncio.gather 并发执行所有 embedding 批次请求:
|
||||
# tasks = []
|
||||
# for i in range(0, len(messages), batch_size):
|
||||
# group_msg = messages[i : i + batch_size]
|
||||
# tasks.append(self.aencode(group_msg))
|
||||
|
||||
# results = await asyncio.gather(*tasks)
|
||||
# for res in results:
|
||||
# data.extend(res)
|
||||
|
||||
# if task_id:
|
||||
# self.embed_state[task_id]["progress"] = len(messages)
|
||||
# self.embed_state[task_id]["status"] = "completed"
|
||||
|
||||
# return data
|
||||
|
||||
for i in range(0, len(messages), batch_size):
|
||||
group_msg = messages[i : i + batch_size]
|
||||
tasks.append(self.aencode(group_msg))
|
||||
|
||||
results = await asyncio.gather(*tasks)
|
||||
for res in results:
|
||||
logger.info(f"Async encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
|
||||
res = await self.aencode(group_msg)
|
||||
data.extend(res)
|
||||
if task_id:
|
||||
self.embed_state[task_id]["progress"] = i + len(group_msg)
|
||||
|
||||
if task_id:
|
||||
self.embed_state[task_id]["progress"] = len(messages)
|
||||
self.embed_state[task_id]["status"] = "completed"
|
||||
|
||||
return data
|
||||
|
||||
Loading…
Reference in New Issue
Block a user