欢迎光临
我们一直在努力

华为OD机试新系统真题 【LLM推理批次最大化】

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]。每个位置同一时间只能分配给一个请求。请选出尽可能多的请求,使它们的区间互不重叠,从而最大化批次的吞吐量。

补充说明:

  • [1, 3] 和 [3, 5] 被认为是重叠。
  • 输入描述

    • 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

      1N105

    • 每个位置

      [

      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

      0LiRi109

    输出描述

    一个整数,表示最多可选多少个不重叠区间的请求。

    样例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;
    }

    赞(0)
    未经允许不得转载:171主机测评 » 华为OD机试新系统真题 【LLM推理批次最大化】
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址