通过率 0% · 提交 0 · 通过 0
给定 n 个请求的 token 长度和 q 个显存预算。为了让同一批尽量容纳更多请求,你可以优先选择长度较短的请求。若选择 c 个请求组成一批,设这 c 个请求中的最大长度为 maxLen,总长度为 sumLen,显存估算为 base + c * fixed + c * maxLen * kv + sumLen * attn。请对每个预算输出最多能放入的请求数。
这类题属于算法机考高频题型中「华为 AI 岗 / 显存估算」方向的高频题型,通常考察对「华为 AI 岗 / 显存估算」的建模能力与边界条件处理。掌握本题的解题思路后,可举一反三应对同类真题方向,稳步提升机考通过率。
第一行输入 n q。第二行输入四个整数 base fixed kv attn。第三行输入 n 个请求长度。随后 q 行,每行一个预算 budget。
输出 q 行,每行一个整数,表示该预算下最多能放入的请求数。
示例 1
输入示例
3 3 50 5 2 1 10 20 30 100 200 1000
输出示例
1 2 3
基础预算查询
示例 2
输入示例
2 2 50 5 2 1 10 20 40 50
输出示例
0 0
预算低于基础显存
时间限制 3000 ms · 内存限制 256 MB
本平台为独立第三方培训机构,与华为技术有限公司无任何关联;课程的服务内容与权益以购买协议为准,学习效果因个人情况而异。「华为 OD」「华为可信」等仅为对岗位与考试方向的客观描述,相关商标归各自权利人所有。
这些是真正决定能不能 AC、但通用题解里常被略过的点。
批处理显存估算 + 多次查询。核心是两个观察:装请求时短的优先不亏,成本关于批大小单调——合起来就是排序、前缀和、二分三件套。
批的成本 base + c·fixed + c·maxLen·kv + sumLen·attn 里,maxLen 和 sumLen 都只会被更长的请求推高。要装下 c 个请求,选最短的 c 个一定成本最低:换任何一个更长的进来,sumLen 变大,maxLen 不降。所以把长度升序排序,「装 c 个」就固定是前 c 个。
排序后设前缀和 pre[c],前 c 个的最大长度就是第 c 个(升序末位)len[c−1]。成本函数:
cost(c) = base + c·fixed + c·len[c−1]·kv + pre[c]·attn
c 增大时每一项都不减,cost 单调不减——对每个预算二分最大的 c 使 cost(c) ≤ budget 即可。q 次查询各一次二分,总复杂度 O(n log n + q log n)。也可以把 cost(1..n) 预先算成数组,对每个预算二分数组,更直观。
c、len、kv、attn 相乘的量级可以冲破 32 位——C++/Java 必须用 64 位整型(long long / long)累加;Python 天然大整数没有这个坑,但换语言重写时最容易在这里翻车。
时间 O(n log n + q log n),空间 O(n)。
样例 1:base=50、fixed=5、kv=2、attn=1,长度升序 10、20、30,前缀和 10、30、60。
预算 100 装下 1 个、200 装下 2 个、1000 装下 3 个。cost 数组单调上升,每个预算在它上面二分位置即可;q 个预算互不影响,别把上一问的选择带进下一问。
显存估算是大模型部署面试的常客,这道题给了一个可复述的框架:一批请求的显存 = 固定开销 + 每请求开销 + KV 缓存(批大小×最长序列×系数)+ 注意力工作区(总 token 数×系数)。追问「为什么 KV 项用 maxLen 而不是各自长度」:批内张量按最长序列对齐填充,短请求也占满对齐后的空间——这正是连续批处理(continuous batching)要解决的浪费。算法侧的可讲点是单调性:成本关于批大小单调,所以「最多装多少」可以二分,这个「先证单调再二分」的动作和二分答案是同一块肌肉。
先手算 cost(1) 和样例对——base、fixed、kv、attn 四项里少乘一项立刻现形;再查排序(没排序时 maxLen 和前缀和全错);大档位输出为负或异常大,直接查 32 位溢出;边界档专看预算小于 cost(1) 时是否输出 0。
# 长度升序排序后装 c 个必取最短的 c 个:成本=base+c·fixed+c·len[c−1]·kv+pre[c]·attn 关于 c 单调;
# 每个预算二分最大 c;C++/Java 用 64 位防溢出。
import bisect
import sys
def solve() -> None:
data = list(map(int, sys.stdin.buffer.read().split()))
if not data:
return
at = 0
n = data[at]; at += 1
q = data[at]; at += 1
base, fixed, kv, attn = data[at:at + 4]; at += 4
lengths = sorted(data[at:at + n]); at += n
prefix = [0]
for value in lengths:
prefix.append(prefix[-1] + value)
need = [base]
for c in range(1, n + 1):
need.append(base + fixed * c + kv * c * lengths[c - 1] + attn * prefix[c])
ans = []
for _ in range(q):
budget = data[at]; at += 1
ans.append(str(max(0, bisect.bisect_right(need, budget) - 1)))
print("\n".join(ans))
if __name__ == "__main__":
solve()
登录后可查看你在本题的历史提交,以及每次的各用例通过情况。