# 通信聚合实现
class CommunicationAggregation:
    def __init__(self, batch_size=32, timeout=0.01):
        self.batch_size = batch_size
        self.timeout = timeout
        self.pending_messages = []
        self.aggregation_timer = None
    
    async def send_message(self, message, target_rank):
        """发送消息（支持聚合）"""
        self.pending_messages.append((message, target_rank))
        
        # 检查是否需要立即发送
        if len(self.pending_messages) >= self.batch_size:
            await self._flush_messages()
        else:
            # 启动聚合定时器
            if self.aggregation_timer is None:
                self.aggregation_timer = asyncio.create_task(
                    self._aggregation_timeout()
                )
    
    async def _aggregation_timeout(self):
        """聚合超时处理"""
        await asyncio.sleep(self.timeout)
        await self._flush_messages()
        self.aggregation_timer = None
    
    async def _flush_messages(self):
        """批量发送消息"""
        if not self.pending_messages:
            return
        
        # 按目标节点分组
        grouped_messages = {}
        for message, target_rank in self.pending_messages:
            if target_rank not in grouped_messages:
                grouped_messages[target_rank] = []
            grouped_messages[target_rank].append(message)
        
        # 批量发送到每个节点
        send_tasks = []
        for target_rank, messages in grouped_messages.items():
            # 将多个消息打包成一个大消息
            batched_message = self._batch_messages(messages)
            task = asyncio.create_task(
                self._send_batched_message(batched_message, target_rank)
            )
            send_tasks.append(task)
        
        await asyncio.gather(*send_tasks)
        self.pending_messages.clear()
    
    def _batch_messages(self, messages):
        """将多个消息打包"""
        return {
            'count': len(messages),
            'data': torch.cat([msg['data'] for msg in messages]),
            'metadata': [msg['metadata'] for msg in messages]
        }
