# 梯度压缩示例
class GradientCompression:
    def __init__(self, compression_ratio=0.1):
        self.compression_ratio = compression_ratio
    
    def compress_gradients(self, gradients):
        """压缩梯度数据"""
        # 只保留最重要的梯度值
        flat_grads = torch.cat([g.flatten() for g in gradients])
        
        # 选择绝对值最大的梯度
        threshold = torch.quantile(torch.abs(flat_grads), 1 - self.compression_ratio)
        mask = torch.abs(flat_grads) > threshold
        
        # 只传输重要的梯度
        compressed_grads = flat_grads[mask]
        compressed_indices = torch.where(mask)[0]
        
        return compressed_grads, compressed_indices
    
    def decompress_gradients(self, compressed_grads, indices, original_shape):
        """解压缩梯度数据"""
        full_grads = torch.zeros(original_shape)
        full_grads[indices] = compressed_grads
        return full_grads
