第一题:混合专家大模型的动态路由与容量掩码
题目描述
在混合专家大模型(MoE)中,为了节省算力,并不是所有的神经网络层都会处理所有的 Token。
当系统包含 E 个专家时,对于输入的每一个 Token,路由器会给各个专家打分,并只将该 Token 分配给得分排名前 K 的专家。同时,为了防止某个专家被分配了过多的 Token 导致显存溢出,系统为每个专家设置了最大处理容量 C。如果某个专家达到了容量上限,后续被路由到该专家的 Token 将被强制"掩蔽"并丢弃。
给定 N 个 Token 对 E 个专家的原始打分矩阵 S(维度 N * E),具体流程如下:
1. Top-K 路由掩码生成:
对于打分矩阵的每一行(代表第 i 个 Token),找出得分排名前 K 的专家索引,生成一个初始的 N * E 二进制路由掩码矩阵 M。
- 如果专家 j 在 Token i 的前 K 名中,则 M[i][j] = 1,否则为 0。
2. 动态容量掩码调度:
按顺序(从 Token 0 到 Token N-1)依次处理每个 Token。维护一个长度为 E 的计数器数组 load,初始全为 0。遍历初始掩码 M,生成最终的生效掩码 M':
- 如果 M[i][j] = 1,且此时 load[j] < C,则该分配生效,M'[i][j] = 1,load[j] 加一。
- 如果 M[i][j] = 1,且此时 load[j] 已经等于 C,则触发容量掩码,强制丢弃,M'[i][j] = 0。
- 对于同一个 Token,被选中的 K 个专家的容量判断是相互独立的,均与当前 Token 处理开始时刻的 load 数组进行比较。
3. 路由惩罚聚合:
根据最终的每个专家的实际负载数组 load,计算当前批次的路由不平衡惩罚值。公式为各专家实际负载的平方和:penalty = sum(load[j]^2)。
输入描述
第 1 行:四个整数,以空格分隔,分别为 Token 数量 N、专家数量 E、每个 Token 激活的专家数 K、单个专家最大容量 C。
接下来 N 行:每行包含 E 个整数,以空格分隔,代表原始打分矩阵 S。
约束条件:1 <= N <= 1000, 1 <= E <= 100, 1 <= K <= E, 1 <= C <= N。
输出描述
第 1 行:一个整数,代表最终的路由惩罚值 penalty。
第 2 行:E 个整数,以空格分隔,代表每个专家最终处理的 Token 数量。
样例 1
输入
4 3 2 21 5 48 1 23 6 52 7 9
输出
91 2 2
说明
Top-K 路由:
- Token 0:得分 [1,5,4],前 2 名是专家 1(5 分)和专家 2(4 分)
- Token 1:得分 [8,1,2],前 2 名是专家 0(8 分)和专家 2(2 分)
- Token 2:得分 [3,6,5],前 2 名是专家 1(6 分)和专家 2(5 分)
- Token 3:得分 [2,7,9],前 2 名是专家 1(7 分)和专家 2(9 分)
容量掩码调度(C = 2):
- Token 0:专家 1、2 容量均充足 → load = [0, 1, 1]
- Token 1:专家 0、2 容量均充足 → load = [1, 1, 2]
- Token 2:专家 1 充足,专家 2 超载(已达 2)→ load = [1, 2, 2]
- Token 3:专家 1 超载(已达 2),专家 2 超载(已达 2)→ load = [1, 2, 2]
惩罚聚合:12 + 22 + 2^2 = 1 + 4 + 4 = 9。
样例 2
输入
4 4 2 25 5 5 12 8 8 91 2 9 97 7 1 1
输出
132 2 1 2
说明
Top-K 路由:
- Token 0:前三名得分同为 5,根据"优先选择索引较小"规则,选专家 0、1
- Token 1:最高分专家 3(9 分),第二高分专家 1 和 2 同为 8 分,优先选索引较小的专家 1
容量掩码调度(C = 2):
- Token 0:专家 0、1 容量均充足 → load = [1, 1, 0, 0]
- Token 1:专家 1、3 容量均充足 → load = [1, 2, 0, 1]
- Token 2:专家 2、3 容量均充足 → load = [1, 2, 1, 2]
- Token 3:专家 0 充足,专家 1 超载 → load = [2, 2, 1, 2]
惩罚聚合:22 + 22 + 12 + 22 = 4 + 4 + 1 + 4 = 13。
解题思路
按照 Token 的输入顺序逐行处理打分矩阵。对于每个 Token,将所有专家按以下规则排序:
排序后的前 K 个专家就是该 Token 的 Top-K 路由结果。随后依次检查这 K 个专家:
- 如果专家 j 当前的负载 load[j] < C,则本次路由生效,load[j] 加一。
一个 Token 的 Top-K 结果中不会重复出现同一个专家。每个专家的容量判断只会修改自己的负载,不会影响该 Token 对其他专家的判断。
不需要显式构造 N * E 的掩码矩阵。对于每个 Token 找到 Top-K 专家后,立即完成容量判断即可。所有 Token 处理完成后,计算 penalty = sum(load[j]^2)。
时间复杂度: O(N * E log E),每个 Token 需要对 E 个专家排序。
空间复杂度: O(N * E),用于存储打分矩阵。排序和负载数组使用 O(E) 额外空间。
更加详细解题思路和CPP、Java代码加我微信获取)
import sys
def compute_routing_penalty(scores, expert_count, top_k, capacity):
"""模拟 MoE 动态路由与容量掩码,计算路由惩罚值"""
# load[j] 表示专家 j 当前已经接收的 Token 数量
load = [0] * expert_count
for token_scores in scores:
# 按得分降序排列,得分相同时按专家索引升序排列
# 排序键为 (-score, index)
ranking = sorted(
range(expert_count),
key=lambda j: (-token_scores[j], j)
)
# 只检查当前 Token 得分排名前 K 的专家
for t in range(top_k):
expert_idx = ranking[t]
# 专家未达到容量上限时,本次路由生效
if load[expert_idx] < capacity:
load[expert_idx] += 1
# 否则触发容量掩码,该 Token 被丢弃(负载不变)
# 计算所有专家实际负载的平方和作为惩罚值
penalty = sum(x * x for x in load)
return penalty, load
def main():
n, expert_count, top_k, capacity = map(
int, sys.stdin.buffer.readline().split()
)
# 读取 N 行专家打分矩阵
scores = []
for _ in range(n):
row = list(map(int, sys.stdin.buffer.readline().split()))
scores.append(row)
penalty, load = compute_routing_penalty(
scores, expert_count, top_k, capacity
)
print(penalty)
print(*load)
if __name__ == "__main__":
main()
第二题:流水线并行阶段划分优化
题目描述
在大规模深度学习训练中,常采用流水线并行(Pipeline Parallelism)来提升训练效率。模型被划分为多个连续阶段(Stage),每个阶段在不同设备上执行。合理的阶段划分需要兼顾计算负载均衡和通信开销最小。
给定一个包含 n 层的模型,需要按顺序划分为 p 个连续阶段。每层有计算时间 time[i],相邻层之间存在通信开销 comm[i]。
如果在层 k 与层 k+1 之间划分阶段,需要产生通信开销 comm[k]。每个阶段的计算时间为该阶段所有层计算时间之和。
所有阶段必须满足:每个阶段的计算时间 <= T(T 为给定的最大阶段计算时间)。
在满足上述约束的情况下,需要选择划分方式,使总通信开销最小。若不存在合法划分方案,输出 -1。
输入描述
一行空格分隔的整数,依次为:
- T:单个阶段允许的最大计算时间,1 <= T <= 100000
- time[1], time[2], ..., time[n]:每层的计算时间,1 <= time[i] <= 100
- comm[1], comm[2], ..., comm[n-1]:相邻层之间的通信开销,1 <= comm[i] <= 10
输出描述
输出一个整数:最小总通信开销。若不存在满足条件的划分方案,输出 -1。
样例 1
输入
5 3 10 2 4 6 3 7 1 1 1 1
输出
2
说明
n=5, p=3, T=10。计算时间:[2, 4, 6, 3, 7],通信开销:[1, 1, 1, 1]。
一种合法划分方式为:[2, 4] | [6, 3] | [7]。
三个阶段计算时间均不超过 T = 10。
划分点在层 2 后和层 4 后,通信开销 = comm[2] + comm[4] = 1 + 1 = 2。
样例 2
输入
4 2 7 3 5 4 2 1 2 3
输出
-1
说明
n=4, p=2, T=7。计算时间:[3, 5, 4, 2],通信开销:[1, 2, 3]。
任意划分都会导致某个阶段计算时间大于 7,因此不存在合法方案。
解题思路
设前缀和 prefix[i] 表示前 i 层的计算时间之和,则层 (k+1) 到层 i 组成一个阶段时,该阶段的计算时间为 prefix[i] - prefix[k]。
定义 dp[s][i] 表示将前 i 层恰好划分为 s 个非空阶段时,所需的最小通信开销。
状态转移:假设最后一个阶段包含第 (k+1) 层到第 i 层,那么前 k 层需要被划分为 s-1 个阶段,并且需要在第 k 层后进行一次阶段划分,产生通信开销 comm[k]。转移方程为:
dp[s][i] = min{ dp[s-1][k] + comm[k] },其中 k 需满足 s-1 <= k < i 且 prefix[i] - prefix[k] <= T。
对于固定的 i,由于每层计算时间均为正数,满足 prefix[i] - prefix[k] <= T 的 k 构成一个连续区间 [left[i], i-1]。
在固定阶段数 s 时,随着 i 增大,合法的 k 区间左右端点均单调右移。因此可以使用单调队列维护滑动区间内 dp[s-1][k] + comm[k] 的最小值,将每层状态转移优化为均摊 O(1)。
实现步骤:
- 使用双指针计算每个 i 的最小合法前置层数 left[i]。
- 从 2 到 p 枚举阶段数量,使用单调队列维护合法切分点对应的最小通信开销。
时间复杂度: O(p * n),每个阶段内每层最多入队出队一次。
空间复杂度: O(n),使用滚动数组 + 前缀和 + 单调队列。
更加详细解题思路和CPP、Java代码加我微信获取:)
# 第2题:流水线并行阶段划分优化
import sys
from collections import deque
def min_communication_cost(layer_count, stage_count, max_time,
compute_times, comm_costs):
"""计算满足各阶段计算时间约束下的最小总通信开销"""
INF = 10 ** 18
# 计算前缀和,用于快速计算任意连续层的总计算时间
prefix = [0] * (layer_count + 1)
for i in range(1, layer_count + 1):
prefix[i] = prefix[i - 1] + compute_times[i]
# left_bound[i] 表示满足 prefix[i] - prefix[k] <= max_time 的最小 k
# 由于所有计算时间均为正数,k 只会单调右移
left_bound = [0] * (layer_count + 1)
k = 0
for i in range(1, layer_count + 1):
while k < i and prefix[i] - prefix[k] > max_time:
k += 1
left_bound[i] = k
# dp[i] 表示将前 i 层恰好划分为当前阶段数时的最小通信开销
# 只有一个阶段时不能产生阶段间通信开销,代价为 0
dp = [INF] * (layer_count + 1)
for i in range(1, layer_count + 1):
if prefix[i] <= max_time:
dp[i] = 0
if stage_count 1:
return -1 if dp[layer_count] INF else dp[layer_count]
# 依次计算划分为 2 到 p 个阶段的结果
for s in range(2, stage_count + 1):
next_dp = [INF] * (layer_count + 1)
# 单调队列:存储 (切分位置 k, dp[k] + comm_costs[k])
# 保持队列中候选值单调递增
mono_queue = deque()
# 至少需要 s 层才能划分为 s 个非空阶段
for i in range(s, layer_count + 1):
# 当最后一个阶段以 i 结尾时,新增的最大切分位置为 i-1
cut_pos = i - 1
# 只有前 cut_pos 层能够被划分为 s-1 个阶段时才可转移
if dp[cut_pos] < INF:
candidate = dp[cut_pos] + comm_costs[cut_pos]
# 保持队列单调递增:弹出队尾所有 >= 当前候选值的元素
# 相同值保留位置更靠后的切分点,使其更晚失效
while mono_queue and mono_queue[-1][1] >= candidate:
mono_queue.pop()
mono_queue.append((cut_pos, candidate))
# 最后一个阶段需要满足计算时间不超过 max_time
min_valid_k = max(s - 1, left_bound[i])
# 删除已经不在合法切分范围内的位置
while mono_queue and mono_queue[0][0] < min_valid_k:
mono_queue.popleft()
# 队首即当前合法范围内的最小通信开销
if mono_queue:
next_dp[i] = mono_queue[0][1]
dp = next_dp
return -1 if dp[layer_count] == INF else dp[layer_count]
def main():
data = list(map(int, sys.stdin.buffer.read().split()))
if not data:
return
layer_count = data[0]
stage_count = data[1]
max_time = data[2]
# 读取每层计算时间(1-indexed 方便与前缀和对齐)
compute_times = [0] * (layer_count + 1)
pos = 3
for i in range(1, layer_count + 1):
compute_times[i] = data[pos]
pos += 1
# 读取相邻层之间的通信开销(1-indexed,comm[k] 表示层 k 与 k+1 之间)
comm_costs = [0] * layer_count
for i in range(1, layer_count):
comm_costs[i] = data[pos]
pos += 1
result = min_communication_cost(
layer_count, stage_count, max_time, compute_times, comm_costs
)
print(result)
if __name__ == "__main__":
main()