LLM推理批次最大化(C++/Py/Java /Js/Go/C)题解
华为OD机试新系统真题 华为OD上机考试新系统真题 8月12号 100分题型
华为OD机试新系统真题目录点击查看: 华为OD机试新系统真题题库目录|机考题库 + 算法考点详解
题目内容
大语言模型推理时,显存中有一个 KV Cache,用于存储各请求的键值对。现有
N
N
N 个推理请求排队,第
i
i
i 个请求需要占用 KV Cache 中一段连续位置
[
L
i
,
R
i
]
[L_i, R_i]
[Li,Ri]。每个位置同一时间只能分配给一个请求。请选出尽可能多的请求,使它们的区间互不重叠,从而最大化批次的吞吐量。
补充说明:
输入描述
- requests:由
N
N
N 个推理请求构成的二维数组,requests[i] 表示第i
i
i 个推理请求,其内容为[
L
i
,
R
i
]
[L_i, R_i]
[Li,Ri]。 - 数组长度
N
N
N 满足1
≤
N
≤
10
5
1 \\le N \\le 10^5
1≤N≤105。 - 每个位置
[
L
i
,
R
i
]
[L_i, R_i]
[Li,Ri] 满足0
≤
L
i
≤
R
i
≤
10
9
0 \\le L_i \\le R_i \\le 10^9
0≤Li≤Ri≤109。
输出描述
一个整数,表示最多可选多少个不重叠区间的请求。
样例1
输入
1,3 2,5 4,7 6,9 8,10 11,12
输出
4
说明 选 [1,3], [4,7], [8,10], [11,12],共 4 个区间互不重叠,为最大可行数。
样例2
输入
1,3 2,4 3,3 4,4
输出
2
说明 选 [1,3], [4,4] 共 2 个不重叠的区间。
样例3
输入
3,5
输出
1
说明 只有 1 个请求,可选择的个数就为 1。
题解
思路:贪心
- 首先将时间区间按照结束时间进行升序排序。
- 从前往后选择区间,当前区间开始时间大于上一个区间的结束时间就选择。
- 对输入区间按照结束时间进行升序排序。
- 定义ans表示选择区间数量,定义lastEnd记录上次选择区间的结束时间。
- 从前往后遍历时间区间{start, end}, 当start > lastEnd时,更新ans++, lastEnd = end
- 最终ans的值就是结果。
本题处理逻辑和前段时间考过的 华为OD新系统机试真题 4.26 -最大化游戏试玩资格分发逻辑基本完全一致。
c++
#include<bits/stdc++.h>
#include <vector>
using namespace std;
// 通用 切割函数 函数 将字符串str根据delimiter进行切割
vector<string> split(const string& str, const string& delimiter) {
vector<string> result;
size_t start = 0;
size_t end = str.find(delimiter);
while (end != string::npos) {
result.push_back(str.substr(start, end – start));
start = end + delimiter.length();
end = str.find(delimiter, start);
}
// 添加最后一个部分
result.push_back(str.substr(start));
return result;
}
int solve(vector<vector<int>>& requests) {
// 按照结束时间进行升序
sort(requests.begin(), requests.end(), [](vector<int>& a, vector<int>& b) {
return a[1] < b[1];
});
int n = requests.size();
int ans = 0;
int lastEnd = –1;
for (int i = 0; i < n; i++) {
int start = requests[i][0];
int end = requests[i][1];
if (start > lastEnd) {
ans++;
lastEnd = end;
}
}
return ans;
}
int main() {
string input;
getline(cin, input);
vector<vector<int>> requests;
vector<string> tmp = split(input, " ");
// 分割字符串获取二维区间数组
for (int i = 0; i < tmp.size(); i++) {
vector<string> tmp1 = split(tmp[i], ",");
requests.push_back({stoi(tmp1[0]), stoi(tmp1[1])});
}
cout << solve(requests);
return 0;
}
Java
import java.util.*;
public class Main {
static int solve(List<int[]> requests) {
// 按照结束时间进行升序
requests.sort((a, b) -> Integer.compare(a[1], b[1]));
int n = requests.size();
int ans = 0;
int lastEnd = –1;
for (int i = 0; i < n; i++) {
int start = requests.get(i)[0];
int end = requests.get(i)[1];
if (start > lastEnd) {
ans++;
lastEnd = end;
}
}
return ans;
}
public static void main(String[] args) {
Scanner scanner = new Scanner(System.in);
String input = scanner.nextLine();
List<int[]> requests = new ArrayList<>();
// 分割字符串获取二维区间数组
String[] tmp = input.split(" ");
for (int i = 0; i < tmp.length; i++) {
String[] tmp1 = tmp[i].split(",");
requests.add(new int[]{
Integer.parseInt(tmp1[0]),
Integer.parseInt(tmp1[1])
});
}
System.out.println(solve(requests));
}
}
Python
# 通用切割函数:Python 自带 split,不需要自定义
def solve(requests):
# 按照结束时间进行升序
requests.sort(key=lambda x: x[1])
n = len(requests)
ans = 0
last_end = –1
for i in range(n):
start = requests[i][0]
end = requests[i][1]
if start > last_end:
ans += 1
last_end = end
return ans
input_str = input()
requests = []
# 分割字符串获取二维区间数组
tmp = input_str.split(" ")
for i in range(len(tmp)):
tmp1 = tmp[i].split(",")
requests.append([int(tmp1[0]), int(tmp1[1])])
print(solve(requests))
JavaScript
function solve(requests) {
// 按照结束时间进行升序
requests.sort((a, b) => a[1] – b[1]);
const n = requests.length;
let ans = 0;
let lastEnd = –1;
for (let i = 0; i < n; i++) {
const start = requests[i][0];
const end = requests[i][1];
if (start > lastEnd) {
ans++;
lastEnd = end;
}
}
return ans;
}
const readline = require("readline");
const rl = readline.createInterface({
input: process.stdin,
output: process.stdout
});
rl.on("line", (input) => {
const requests = [];
// 分割字符串获取二维区间数组
const tmp = input.split(" ");
for (let i = 0; i < tmp.length; i++) {
const tmp1 = tmp[i].split(",");
requests.push([
Number(tmp1[0]),
Number(tmp1[1])
]);
}
console.log(solve(requests));
rl.close();
});
Go
package main
import (
"bufio"
"fmt"
"os"
"sort"
"strconv"
"strings"
)
func solve(requests [][]int) int {
// 按照结束时间进行升序
sort.Slice(requests, func(i, j int) bool {
return requests[i][1] < requests[j][1]
})
n := len(requests)
ans := 0
lastEnd := –1
for i := 0; i < n; i++ {
start := requests[i][0]
end := requests[i][1]
if start > lastEnd {
ans++
lastEnd = end
}
}
return ans
}
func main() {
reader := bufio.NewReader(os.Stdin)
input, _ := reader.ReadString('\\n')
input = strings.TrimSpace(input)
var requests [][]int
// 分割字符串获取二维区间数组
tmp := strings.Split(input, " ")
for i := 0; i < len(tmp); i++ {
tmp1 := strings.Split(tmp[i], ",")
start, _ := strconv.Atoi(tmp1[0])
end, _ := strconv.Atoi(tmp1[1])
requests = append(requests, []int{start, end})
}
fmt.Println(solve(requests))
}
C语言
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
typedef struct {
int start;
int end;
} Request;
// 按照结束时间进行升序
int compare(const void* a, const void* b) {
Request* x = (Request*)a;
Request* y = (Request*)b;
return x->end – y->end;
}
int solve(Request requests[], int n) {
// 按照结束时间进行升序
qsort(requests, n, sizeof(Request), compare);
int ans = 0;
int lastEnd = –1;
for (int i = 0; i < n; i++) {
int start = requests[i].start;
int end = requests[i].end;
if (start > lastEnd) {
ans++;
lastEnd = end;
}
}
return ans;
}
int main() {
char input[10000009];
Request requests[100005];
int n = 0;
fgets(input, sizeof(input), stdin);
input[strcspn(input, "\\n")] = '\\0';
// 分割字符串获取二维区间数组
char* token = strtok(input, " ");
while (token != NULL) {
int start, end;
sscanf(token, "%d,%d", &start, &end);
requests[n].start = start;
requests[n].end = end;
n++;
token = strtok(NULL, " ");
}
printf("%d", solve(requests, n));
return 0;
}

