多多的GPU批处理调度
拼多多技术岗 7月19号笔试 第二题
题目内容
多多是一个大模型架构师,在部署大语言模型时为了提高 GPUGPUGPU 的利用率,推理引擎通常会将多个用户的请求合并成一个批次进行并行计算。同一批次内如果请求长度不同,短请求会被填充空白 TokenTokenToken 补到该批次最长请求的长度。
为了在"高并发"和"少浪费"之间取得平衡,多多制定了如下的批处理调度策略:
- 显存限制:每个批次最多只能包含 CCC 个请求。
- 算力浪费限制:同一个批次内,最长请求的 TokenTokenToken 数与最短请求的 TokenTokenToken 数之差不能超过 KKK。
- 请求的执行顺序可以被任意打乱。
在实际业务中,偶尔会出现极个别长度异常的"离群"请求。为了不让这些离群点拖累整体吞吐量,系统引入了降级策略:推理引擎最多允许直接丢弃 MMM 个请求,被丢弃的请求可以不处理。
现在,服务器积压了 NNN 个请求,第 iii 个请求包含 LiL_iLi 个Token。请帮多多设计一个调度算法,在最多丢弃 MMM 个请求的前提下,最少需要将剩下的请求划分成多少个批次?
输入描述
第一行包含一个整数 T(1≤T≤3)T(1 \le T \le 3)T(1≤T≤3),表示测试用例的数量。
对于每个测试用例:
第一行包含四个整数 N, M, K, CN,\ M,\ K,\ CN, M, K, C,分别表示请求总数、最多允许丢弃的请求数、允许的最大长度差、单个批次最大请求数。
1≤N≤105, 0≤M≤50, 0≤K≤109, 1≤C≤N1 \le N \le 10^5,\ 0 \le M \le 50,\ 0 \le K \le 10^9,\ 1 \le C \le N1≤N≤105, 0≤M≤50, 0≤K≤109, 1≤C≤N
第二行包含 NNN 个整数 L1,L2,...,LnL_1,L_2,\dots,L_nL1,L2,...,Ln,表示每个请求的 TokenTokenToken 长度。
1≤L≤1091 \le L \le 10^91≤L≤109
输出描述
对于每个测试用例,输出一行包含一个整数,表示最少需要的批次数量。
样例1
输入
2
5 1 1 2
10 11 15 16 17
4 2 10 3
1 100 2 200
输出
2
1
说明
第一个样例 NNN=555,MMM=111,KKK=111,CCC=222,长度为 10,11,15,16,1710, 11, 15, 16, 1710,11,15,16,17。
如果不丢弃任何请求,至少需要 333 个批次(例如 (10,11),(15,16),(17)(10, 11), (15, 16), (17)(10,11),(15,16),(17))。
如果允许丢弃 111 个请求,我们可以选择丢弃 151515,剩下的 (10,11)(10, 11)(10,11) 编为一批,最大差值为 111;(16,17)(16, 17)(16,17) 编为一批,最大差值为 111,此时只需要 222 个批次。
第二个样例 NNN=444,MMM=222,KKK=101010,CCC=333,长度为 1,100,2,2001, 100, 2, 2001,100,2,200。
丢弃离群点 100100100 和 200200200,剩下 (1,2)(1, 2)(1,2) 可以放在一个批次中,因此最少需要 111 个批次。
题解
思路
解题算法:动态规划
- 由于题目允许任意调整请求顺序,所以真正影响能否放到同一个 batch 的只有最长请求和最短请求,因此先对数据进行升序排序。
- 使用双指针预处理
A[i],如果第 i 个请求作为某个 batch 中最大的请求,那么前面至少要处理到哪个位置。限制来自于两个条件- 长度差限制,最大值和最小值差
- 容量限制C
- 动态规划设置状态
f[i][del]表示前 i 个请求,删除 del 个请求,最少需要多少个 batch。对于每个请求有以下几种选择- 删除第i个请求:
f[i][del]=f[i-1][del-1] - 保留第i个请求,那么它一定属于最后一个 batch,并且是最后一个 batch 的最大请求。所以
f[i][del]=min(f[q][del])+1其中A[i]<=q<=i-1
- 删除第i个请求:
- 从第三步为了加速处理计算
f[i][del]=min(f[q][del])+1其中A[i]<=q<=i-1,引入h[i][del]表示min(f[i][del],f[i-1][del-1],f[i-2][del-2],...)前 i 个请求已经处理完,允许把若干个请求删掉之后得到的最优答案. 对于固定删除数量del,通过单调递增求最小值。 - 下面总体时间复杂度为
O(NM)
C++
cpp
#include<bits/stdc++.h>
using namespace std;
using ll = long long;
const ll INF = 4e18;
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int T;
cin >> T;
while (T--) {
int N,M,K,C;
cin >> N >> M >> K >> C;
vector<ll> a(N + 1);
for (int i = 1; i <= N; i++) {
cin >> a[i];
}
// 不必考虑顺序,需要考虑最大值和最小值差异,直接进行排序
sort(a.begin() + 1, a.end());
// 最多丢弃数量
M = min(M, N);
// 如果第 i 个请求作为最后一个batch的最大值, 那么那么前面必须处理到的位置
vector<int> A(N + 1);
int left = 1;
for (int i = 1; i <= N; i++) {
while (left <= i && a[i] - a[left] > K) {
left++;
}
A[i] = max(left - 1, i - C);
}
// f[i][j] 处理完i删除j各元素,最小批次数量的滚动数组
vector<ll> fPrev(N + 1, INF);
vector<ll> hPrev(N + 1, INF);
vector<ll> fCur(N + 1);
vector<ll> hCur(N + 1);
ll ans = INF;
// 枚举删除数量
for (int del = 0; del <= M; del++) {
fill(fCur.begin(), fCur.end(), INF);
fill(hCur.begin(), hCur.end(), INF);
// 当前删除数量为0
if (del == 0) {
fCur[0] = 0;
hCur[0] = 0;
}
// 单调队列维护 d[q][del]
deque<pair<int,ll>> dq;
for (int i = 1; i <= N; i++) {
if (hCur[i - 1] < INF) {
while (!dq.empty() && dq.back().second >= hCur[i-1]) {
dq.pop_back();
}
dq.push_back({i - 1, hCur[i-1]});
}
// 删除窗口外元素
while (!dq.empty() && dq.front().first < A[i]) {
dq.pop_front();
}
// 保留第i个请求
if (!dq.empty()) {
fCur[i] = dq.front().second + 1;
}
// 删除第i个请求
if (del > 0) {
fCur[i] = min(fCur[i], fPrev[i-1]);
}
// 更新h
if (del > 0) {
hCur[i] = min(fCur[i], hPrev[i-1]);
} else {
hCur[i] = fCur[i];
}
}
ans = min(ans, fCur[N]);
fPrev.swap(fCur);
hPrev.swap(hCur);
}
cout << ans << endl;
}
return 0;
}
java
java
import java.io.*;
import java.util.*;
public class Main {
static final long INF = (long) 4e18;
public static void main(String[] args) throws Exception {
BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
int idx = 0;
List<String> tokens = new ArrayList<>();
String line;
while ((line = br.readLine()) != null) {
if (!line.trim().isEmpty()) {
tokens.addAll(Arrays.asList(line.trim().split("\\s+")));
}
}
int p = 0;
int T = Integer.parseInt(tokens.get(p++));
while (T-- > 0) {
int N = Integer.parseInt(tokens.get(p++));
int M = Integer.parseInt(tokens.get(p++));
long K = Long.parseLong(tokens.get(p++));
int C = Integer.parseInt(tokens.get(p++));
long[] a = new long[N + 1];
for (int i = 1; i <= N; i++) {
a[i] = Long.parseLong(tokens.get(p++));
}
// 不必考虑顺序,需要考虑最大值和最小值差异,直接进行排序
Arrays.sort(a, 1, N + 1);
// 最多丢弃数量
M = Math.min(M, N);
// 如果第 i 个请求作为最后一个batch的最大值, 那么那么前面必须处理到的位置
int[] A = new int[N + 1];
int left = 1;
for (int i = 1; i <= N; i++) {
while (left <= i && a[i] - a[left] > K) {
left++;
}
A[i] = Math.max(left - 1, i - C);
}
// f[i][j] 处理完i删除j各元素,最小批次数量的滚动数组
long[] fPrev = new long[N + 1];
long[] hPrev = new long[N + 1];
long[] fCur = new long[N + 1];
long[] hCur = new long[N + 1];
Arrays.fill(fPrev, INF);
Arrays.fill(hPrev, INF);
long ans = INF;
// 枚举删除数量
for (int del = 0; del <= M; del++) {
Arrays.fill(fCur, INF);
Arrays.fill(hCur, INF);
// 当前删除数量为0
if (del == 0) {
fCur[0] = 0;
hCur[0] = 0;
}
// 单调队列维护 h[q][del]
ArrayDeque<Integer> deque = new ArrayDeque<>();
for (int i = 1; i <= N; i++) {
if (hCur[i - 1] < INF) {
while (!deque.isEmpty() && hCur[deque.peekLast()] >= hCur[i - 1]) {
deque.pollLast();
}
deque.offerLast(i - 1);
}
// 删除窗口外元素
while (!deque.isEmpty() && deque.peekFirst() < A[i]) {
deque.pollFirst();
}
// 保留第i个请求
if (!deque.isEmpty()) {
fCur[i] = hCur[deque.peekFirst()] + 1;
}
// 删除第i个请求
if (del > 0) {
fCur[i] = Math.min(fCur[i], fPrev[i - 1]);
}
// 更新h
if (del > 0) {
hCur[i] = Math.min(fCur[i], hPrev[i - 1]);
} else {
hCur[i] = fCur[i];
}
}
ans = Math.min(ans, fCur[N]);
long[] tmp = fPrev;
fPrev = fCur;
fCur = tmp;
tmp = hPrev;
hPrev = hCur;
hCur = tmp;
}
System.out.println(ans);
}
}
}
python
python
from collections import deque
INF = 4 * 10 ** 18
data = input().split()
while True:
try:
data += input().split()
except:
break
ptr = 0
T = int(data[ptr])
ptr += 1
for _ in range(T):
N = int(data[ptr])
ptr += 1
M = int(data[ptr])
ptr += 1
K = int(data[ptr])
ptr += 1
C = int(data[ptr])
ptr += 1
a = [0] * (N + 1)
for i in range(1, N + 1):
a[i] = int(data[ptr])
ptr += 1
# 不必考虑顺序,需要考虑最大值和最小值差异,直接进行排序
a[1:] = sorted(a[1:])
# 最多丢弃数量
M = min(M, N)
# 如果第 i 个请求作为最后一个batch的最大值, 那么那么前面必须处理到的位置
A = [0] * (N + 1)
left = 1
for i in range(1, N + 1):
while left <= i and a[i] - a[left] > K:
left += 1
A[i] = max(left - 1, i - C)
# f[i][j] 处理完i删除j各元素,最小批次数量的滚动数组
fPrev = [INF] * (N + 1)
hPrev = [INF] * (N + 1)
ans = INF
# 枚举删除数量
for delete in range(M + 1):
fCur = [INF] * (N + 1)
hCur = [INF] * (N + 1)
# 当前删除数量为0
if delete == 0:
fCur[0] = 0
hCur[0] = 0
# 单调队列维护 h[q][del]
dq = deque()
for i in range(1, N + 1):
if hCur[i - 1] < INF:
while dq and hCur[dq[-1]] >= hCur[i - 1]:
dq.pop()
dq.append(i - 1)
# 删除窗口外元素
while dq and dq[0] < A[i]:
dq.popleft()
# 保留第i个请求
if dq:
fCur[i] = hCur[dq[0]] + 1
# 删除第i个请求
if delete > 0:
fCur[i] = min(fCur[i], fPrev[i - 1])
# 更新h
if delete > 0:
hCur[i] = min(fCur[i], hPrev[i - 1])
else:
hCur[i] = fCur[i]
ans = min(ans, fCur[N])
fPrev = fCur
hPrev = hCur
print(ans)
javascript
js
const readline = require("readline");
const rl = readline.createInterface({
input: process.stdin,
output: process.stdout
});
const input = [];
rl.on("line", line => {
input.push(...line.trim().split(/\s+/));
});
rl.on("close", () => {
let idx = 0;
const T = Number(input[idx++]);
const INF = 4000000000000000000n;
const ans = [];
for (let tc = 0; tc < T; tc++) {
const N = Number(input[idx++]);
let M = Number(input[idx++]);
const K = BigInt(input[idx++]);
const C = Number(input[idx++]);
const a = new Array(N + 1).fill(0n);
for (let i = 1; i <= N; i++) {
a[i] = BigInt(input[idx++]);
}
// 不必考虑顺序,需要考虑最大值和最小值差异,直接进行排序
const arr = a.slice(1).sort((x, y) => (x < y ? -1 : 1));
for (let i = 1; i <= N; i++) {
a[i] = arr[i - 1];
}
// 最多丢弃数量
M = Math.min(M, N);
// 如果第 i 个请求作为最后一个batch的最大值, 那么那么前面必须处理到的位置
const A = new Array(N + 1).fill(0);
let left = 1;
for (let i = 1; i <= N; i++) {
while (left <= i && a[i] - a[left] > K) {
left++;
}
A[i] = Math.max(left - 1, i - C);
}
// f[i][j] 处理完i删除j各元素,最小批次数量的滚动数组
let fPrev = new Array(N + 1).fill(INF);
let hPrev = new Array(N + 1).fill(INF);
let answer = INF;
// 枚举删除数量
for (let del = 0; del <= M; del++) {
const fCur = new Array(N + 1).fill(INF);
const hCur = new Array(N + 1).fill(INF);
// 当前删除数量为0
if (del === 0) {
fCur[0] = 0n;
hCur[0] = 0n;
}
// 单调队列维护 h[q][del]
const dq = [];
let head = 0;
for (let i = 1; i <= N; i++) {
if (hCur[i - 1] < INF) {
while (dq.length > head && hCur[dq[dq.length - 1]] >= hCur[i - 1]) {
dq.pop();
}
dq.push(i - 1);
}
// 删除窗口外元素
while (dq.length > head && dq[head] < A[i]) {
head++;
}
// 保留第i个请求
if (dq.length > head) {
fCur[i] = hCur[dq[head]] + 1n;
}
// 删除第i个请求
if (del > 0) {
if (fPrev[i - 1] < fCur[i]) {
fCur[i] = fPrev[i - 1];
}
}
// 更新h
if (del > 0) {
hCur[i] = fCur[i] < hPrev[i - 1] ? fCur[i] : hPrev[i - 1];
} else {
hCur[i] = fCur[i];
}
}
if (fCur[N] < answer) {
answer = fCur[N];
}
fPrev = fCur;
hPrev = hCur;
}
ans.push(answer.toString());
}
console.log(ans.join("\n"));
});
Go
go
package main
import (
"bufio"
"fmt"
"os"
"sort"
)
const INF int64 = 4e18
type Pair struct {
idx int
val int64
}
func main() {
in := bufio.NewReader(os.Stdin)
var T int
fmt.Fscan(in, &T)
for ; T > 0; T-- {
var N, M, C int
var K int64
fmt.Fscan(in, &N, &M, &K, &C)
a := make([]int64, N+1)
for i := 1; i <= N; i++ {
fmt.Fscan(in, &a[i])
}
// 不必考虑顺序,需要考虑最大值和最小值差异,直接进行排序
sort.Slice(a[1:], func(i, j int) bool {
return a[i+1] < a[j+1]
})
// 最多丢弃数量
if M > N {
M = N
}
// 如果第 i 个请求作为最后一个batch的最大值, 那么那么前面必须处理到的位置
A := make([]int, N+1)
left := 1
for i := 1; i <= N; i++ {
for left <= i && a[i]-a[left] > K {
left++
}
if left-1 > i-C {
A[i] = left - 1
} else {
A[i] = i - C
}
}
// f[i][j] 处理完i删除j各元素,最小批次数量的滚动数组
fPrev := make([]int64, N+1)
hPrev := make([]int64, N+1)
for i := 0; i <= N; i++ {
fPrev[i] = INF
hPrev[i] = INF
}
ans := INF
// 枚举删除数量
for del := 0; del <= M; del++ {
fCur := make([]int64, N+1)
hCur := make([]int64, N+1)
for i := 0; i <= N; i++ {
fCur[i] = INF
hCur[i] = INF
}
// 当前删除数量为0
if del == 0 {
fCur[0] = 0
hCur[0] = 0
}
// 单调队列维护 h[q][del]
dq := make([]Pair, 0)
head := 0
for i := 1; i <= N; i++ {
if hCur[i-1] < INF {
for len(dq) > head && dq[len(dq)-1].val >= hCur[i-1] {
dq = dq[:len(dq)-1]
}
dq = append(dq, Pair{i - 1, hCur[i-1]})
}
// 删除窗口外元素
for len(dq) > head && dq[head].idx < A[i] {
head++
}
// 保留第i个请求
if len(dq) > head {
fCur[i] = dq[head].val + 1
}
// 删除第i个请求
if del > 0 && fPrev[i-1] < fCur[i] {
fCur[i] = fPrev[i-1]
}
// 更新h
if del > 0 {
if fCur[i] < hPrev[i-1] {
hCur[i] = fCur[i]
} else {
hCur[i] = hPrev[i-1]
}
} else {
hCur[i] = fCur[i]
}
}
if fCur[N] < ans {
ans = fCur[N]
}
fPrev = fCur
hPrev = hCur
}
fmt.Println(ans)
}
}