解题思路
核心思路
每个梯度任务都要分到恰好一张卡上。直接枚举:第 i 个任务依次尝试 k 张卡,用数组记下各卡当前耗时。如果某张卡加上这个任务后,耗时已经大于目前找到的最小「最大耗时」,这一支就不用再搜。搜完得到最小的 max_time,以及一组达到它的各卡负载,再按
cost=j∑sum_timej⋅base_lr⋅decay_ratemax_time−sum_timej
题目描述
在大模型分布式训练场景下,要把若干梯度更新任务分配到多张 NPU 上。目标是让各卡的「梯度计算总耗时」尽量均衡(即最小化「最大卡总耗时」),再按学习率衰减规则汇总总训练成本。
规则如下:
- 任务分配:将 grad_times 划分成 k 个子集(一张 NPU 对应一个子集)。每个任务必须分到恰好一个子集;在所有划分中,取「子集总和的最大值」最小的那些划分。
- 成本计算:记 max_time 为各卡总耗时的最大值,sum_time_j 为第 j 张卡的总耗时。则
- 第 j 张卡学习率:lr_j=base_lr×decay_ratemax_time−sum_time_j;
- 第 j 张卡训练成本:sum_time_j×lr_j;
- 总训练成本:所有卡成本之和。
输入描述
一行字符串,用 ; 分成四段:梯度耗时列表;NPU数量;基础学习率;衰减系数。例如:8,5,4,3,3,2,1;3;0.1;0.8。
- 梯度耗时列表(grad_times):逗号分隔的一维数组,每个元素为单个梯度更新任务的耗时(非负整数);
- NPU 数量(k):正整数,满足 1≤k≤len(grad_times);
- 基础学习率(base_lr):浮点数;
- 衰减系数(decay_rate):浮点数,满足 0<decay_rate<1。
输出描述
一行字符串:最小最大总耗时,总训练成本。例如:9,2.4400。
- 最小最大总耗时:整数;
- 总训练成本:保留 4 位小数的浮点数。
样例1
输入
8,5,4,3,3,2,1;3;0.1;0.8
输出
9,2.4400
说明
解析得 grad_times=[8,5,4,3,3,2,1],k=3,base_lr=0.1,decay_rate=0.8。
总耗时 8+5+4+3+3+2+1=26,理想均分约 26/3≈8.67,且最大单任务为 8,故最小可能的最大子集和至少为 9。
一组达到该下界的划分:
- NPU1:8+1=9
- NPU2:5+4=9
- NPU3:3+3+2=8
负载为 9,9,8,max_time=9。
成本:
- 负载 9:9×0.1×0.80=0.9(两张)
- 负载 8:8×0.1×0.81=0.64
总成本 0.9+0.9+0.64=2.44,输出 9,2.4400。
样例2
输入
6;1;0.2;0.9
输出
6,1.2000
说明
仅 1 张 NPU、一个任务,故 max_time=6,sum_time=6。
lr=0.2×0.90=0.2,总成本 6×0.2=1.2,输出 6,1.2000。