拼多多笔试真题-多多的GPU批处理调度(C++/Py/Java /Js/Go)

多多的GPU批处理调度

拼多多技术岗 7月19号笔试 第二题

题目内容

多多是一个大模型架构师,在部署大语言模型时为了提高 GPUGPUGPU 的利用率,推理引擎通常会将多个用户的请求合并成一个批次进行并行计算。同一批次内如果请求长度不同,短请求会被填充空白 TokenTokenToken 补到该批次最长请求的长度。

为了在"高并发"和"少浪费"之间取得平衡,多多制定了如下的批处理调度策略:

  1. 显存限制:每个批次最多只能包含 CCC 个请求。
  2. 算力浪费限制:同一个批次内,最长请求的 TokenTokenToken 数与最短请求的 TokenTokenToken 数之差不能超过 KKK。
  3. 请求的执行顺序可以被任意打乱。
    在实际业务中,偶尔会出现极个别长度异常的"离群"请求。为了不让这些离群点拖累整体吞吐量,系统引入了降级策略:推理引擎最多允许直接丢弃 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 个批次。

题解

思路

解题算法:动态规划

  1. 由于题目允许任意调整请求顺序,所以真正影响能否放到同一个 batch 的只有最长请求和最短请求,因此先对数据进行升序排序。
  2. 使用双指针预处理A[i],如果第 i 个请求作为某个 batch 中最大的请求,那么前面至少要处理到哪个位置。限制来自于两个条件
    • 长度差限制,最大值和最小值差
    • 容量限制C
  3. 动态规划设置状态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
  4. 从第三步为了加速处理计算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,通过单调递增求最小值。
  5. 下面总体时间复杂度为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)
	}
}
相关推荐
无限码力3 天前
拼多多笔试真题-多多的灰度发布(C++/Py/Java /Js/Go)
拼多多·拼多多机试·拼多多技术岗笔试·拼多多笔试·拼多多技术岗笔试真题
市象19 天前
新拼姆会是大号版“SHEIN”吗?
拼多多
无限码力22 天前
拼多多笔试真题【多多的特殊三元组】
拼多多·拼多多笔试真题·拼多多笔试题库·拼多多技术岗笔试·pdd笔试笔试真题
无限码力22 天前
拼多多笔试真题-多多的Boss挑战(C++/Py/Java /Js/Go)
拼多多·拼多多笔试真题·拼多多技术岗笔试题目·拼多多机试·pdd笔试真题
无限码力24 天前
拼多多笔试真题-多多捕蝇(C++/Py/Java /Js/Go)
拼多多·拼多多笔试真题·拼多多技术岗笔试题目·拼多多机试·pdd笔试真题
无限码力1 个月前
拼多多笔试真题-多多的营救行动(C++/Py/Java /Js/Go)
拼多多·拼多多笔试真题·拼多多笔试题库·拼多多技术岗笔试
无限码力1 个月前
拼多多笔试真题-对角线遍历矩阵(C++/Py/Java /Js/Go)
矩阵·拼多多·拼多多笔试真题·拼多多技术岗笔试题目·拼多多机试
27669582921 年前
拼多多小程序 anti_content 分析
java·javascript·python·node·拼多多·anti-content·anti_content
27669582921 年前
拼多多 anti-token unidbg 分析
java·python·go·拼多多·pdd·pxx·anti-token