通过率 0% · 提交 0 · 通过 0
给定 n 个带标签的样本和 m 个查询点,对每个查询点:计算它到每个样本的平方欧氏距离(不做开方,特征为整数,距离为精确整数),取距离最小的 k 个样本作为近邻;距离相同时样本编号(输入顺序,从 0 开始)小者优先。在 k 个近邻的标签中做多数表决,票数相同时类别编号小者当选。输出每个查询点的预测类别。
这类题属于算法机考高频题型中「华为 AI 岗 / 排序」方向的高频题型,通常考察对「华为 AI 岗 / 排序」的建模能力与边界条件处理。掌握本题的解题思路后,可举一反三应对同类真题方向,稳步提升机考通过率。
第一行输入 n d k m,分别为样本数、特征维数、近邻数和查询数。接下来 n 行,每行 d 个整数特征和一个整数标签。随后 m 行,每行 d 个整数,表示查询点。
输出 m 行,每行一个整数,表示对应查询点的预测类别。
示例 1
输入示例
5 2 3 2 0 0 0 1 0 0 5 5 1 6 5 1 0 1 0 0 0 5 4
输出示例
0 1
三近邻多数表决。
示例 2
输入示例
3 1 1 1 -2 5 2 9 4 9 0
输出示例
5
到 -2 和 2 距离都是 4,取样本编号更小的 0 号,类别 5。
时间限制 2000 ms · 内存限制 256 MB
本平台为独立第三方培训机构,与华为技术有限公司无任何关联;课程的服务内容与权益以购买协议为准,学习效果因个人情况而异。「华为 OD」「华为可信」等仅为对岗位与考试方向的客观描述,相关商标归各自权利人所有。
这些是真正决定能不能 AC、但通用题解里常被略过的点。
KNN 分类的精确模拟。特征全是整数、距离不开方,整个计算过程没有一个浮点数——考点全在两级并列规则和 64 位整型上。
平方欧氏距离 Σ(x_i − s_i)²,整数运算精确无误差,这正是不开方的好处(开方引入浮点,比较就要容差,规则复杂度翻倍)。
取前 k 个近邻的规则:距离小者优先,距离相同时样本编号小者优先。把每个样本表示成 (距离, 编号) 二元组排序取前 k,规则自动满足。n ≤ 2000、m ≤ 200,O(m·n log n) 全排序足够;不必上堆或快速选择。
k 个近邻的标签计票,票数最多者当选;票数相同时类别编号小者当选。和距离并列一样用复合键解决:min(候选类别, key=(−票数, 类别))。两级并列各有专门用例,凭默认顺序(字典插入序、稳定排序的巧合)能过样例过不了全集。
单维差最大 20000,平方 4×10⁸,8 维累加可到 3.2×10⁹——超出 32 位 int。C++ 用 long long,Java 用 long;Python 天然大整数无此坑,但换语言重写时这是第一翻车点。
时间 O(m·n·(d + log n)),空间 O(n)。
样例 1:五个样本 (0,0)、(1,0)、(5,5)、(6,5)、(0,1),标签 0、0、1、1、0,k=3。
查询 (0,0):平方距离 0、1、50、61、1。样本 #1 和 #4 距离并列为 1,按编号 #1 优先——排序键 (距离, 编号) 自动处理。近邻 {#0, #1, #4},标签 0、0、0,输出 0。
查询 (5,4):距离 41、32、1、2、34,近邻 {#2, #3, #1},标签 1、1、0,表决 2:1 输出 1。整个过程没有一个浮点数出现——这是不开方设计的全部意义。
KNN 的面试价值在于它是「无参数模型」的代表:fit 只是把训练集存下来,全部计算发生在预测时。标准追问链条:k 太小容易被噪声点带偏(方差大),k 太大把远处样本也拉进来投票(偏差大),k 用验证集选;预测代价 O(n·d) 随训练集线性增长,大规模场景要靠近似最近邻(向量检索)救场——这句话正好把这道题接到向量数据库的语境上。实现层面的可讲点就是两级并列规则:距离并列按编号、票数并列按类别,规则显式化才能让结果可复现。
先用「两个样本距离相同」的小用例验证编号优先;再用「两类票数相同」的用例验证类别取小;C++/Java 在大坐标档位错、小档位对,基本就是 32 位溢出。全对只错一两个查询时,重点查 k 个近邻的边界——排序后取前 k 是否含并列处理。
工程小注:排序取前 k 也可以换成 heapq.nsmallest(k, ...),复杂度从 O(n log n) 降到 O(n log k)——规模内两者都过,但把 (距离, 编号) 元组直接交给堆,并列规则依旧自动成立,这个「复合键交给容器」的手法与前面 Top K、调度类题目一脉相承。
# 整数平方距离(不开方,精确比较),近邻按 (距离,编号) 排序取前 k;
# 表决键=(−票数,类别),两级并列都取小;C++/Java 距离用 64 位。
import sys
def solve() -> None:
data = sys.stdin.buffer.read().split()
if not data:
return
it = iter(data)
n, d, k, m = (int(next(it)) for _ in range(4))
x, y = [], []
for _ in range(n):
x.append([int(next(it)) for _ in range(d)])
y.append(int(next(it)))
queries = [[int(next(it)) for _ in range(d)] for _ in range(m)]
out = []
for q in queries:
order = []
for idx, row in enumerate(x):
dist = 0
for j in range(d):
diff = row[j] - q[j]
dist += diff * diff
order.append((dist, idx))
order.sort()
votes = {}
for dist, idx in order[:k]:
votes[y[idx]] = votes.get(y[idx], 0) + 1
best_label = -1
best_count = -1
for label, count in votes.items():
if count > best_count or (count == best_count and label < best_label):
best_label = label
best_count = count
out.append(str(best_label))
print("\n".join(out))
if __name__ == "__main__":
solve()
登录后可查看你在本题的历史提交,以及每次的各用例通过情况。
© 2026 广州慕课网络科技有限公司 · 吴师兄学算法官网 版权所有