通过率 0% · 提交 0 · 通过 0
给定若干棵二叉决策树和一批样本,请模拟树模型推理。内部节点根据某个特征是否小于等于阈值决定走左子树或右子树;叶子节点给出类别。每棵树得到一个类别后,用多数表决得到最终结果;如果多个类别票数相同,选择数值更小的类别。
这类题属于算法机考高频题型中「华为 AI 岗 / 决策树」方向的高频题型,通常考察对「华为 AI 岗 / 决策树」的建模能力与边界条件处理。掌握本题的解题思路后,可举一反三应对同类真题方向,稳步提升机考通过率。
第一行输入 T d q,表示树的数量、特征维度和待预测样本数。随后对每棵树,先输入节点数 s,再输入 s 行节点。节点编号从 1 开始,1 号为根。叶子节点格式为 L label;内部节点格式为 N feature threshold left right。最后输入 q 行样本,每行 d 个特征值。
输出 q 行,每行一个整数,表示对应样本的最终类别。
示例 1
输入示例
1 1 3 3 N 1 5 2 3 L 0 L 1 3 5 6
输出示例
0 0 1
单棵树基础推理
示例 2
输入示例
3 2 3 3 N 1 5 2 3 L 0 L 1 3 N 2 0 2 3 L 1 L 2 3 N 1 2 2 3 L 0 L 2 1 -1 6 1 3 5
输出示例
0 2 2
多棵树多数表决
时间限制 3000 ms · 内存限制 256 MB
本平台为独立第三方培训机构,与华为技术有限公司无任何关联;课程的服务内容与权益以购买协议为准,学习效果因个人情况而异。「华为 OD」「华为可信」等仅为对岗位与考试方向的客观描述,相关商标归各自权利人所有。
这些是真正决定能不能 AC、但通用题解里常被略过的点。
多棵决策树的推理模拟 + 多数表决。没有训练过程,纯粹考「把规则走对」:节点怎么走、平票怎么裁、深树怎么不炸。
每棵树用数组存节点(编号从 1 开始,1 号是根)。内部节点 N feature threshold left right:样本第 feature 个特征小于等于阈值走左,否则走右——注意是 <=,阈值恰好相等的用例专门卡这个方向。叶子节点 L label 直接给类别。
用循环走树,不用递归。 节点总数上限 5000,链状树深度可以到几千,Python 默认递归上限 1000 层,递归版会在深树用例上 RecursionError。循环版从根出发 while 到叶子即可,天然免疫:
node = 1
while nodes[node][0] == 'N':
_, f, th, l, r = nodes[node]
node = l if x[f] <= th else r
label = nodes[node][1]T 棵树各出一个类别,统计票数取最多;票数相同时取数值更小的类别。用 min(votes, key=lambda c: (-votes[c], c)) 一行写清,或者排序时用 (−票数, 类别) 作键。平票用例是标配,别依赖字典遍历顺序。
1. 逐棵树读入节点表(注意每棵树节点数不同,读入指针别串位)。 2. 对每个样本走完全部树收集类别,表决输出。
时间 O(q·T·h),h 为树的最大深度;空间 O(总节点数 + d)。
<= 写成 <,阈值相等的样本走错方向。样例 1 只有一棵树三个节点:根 N 1 5 2 3(第 1 个特征与阈值 5 比较),左右孩子是叶子 L 0 和 L 1。三个样本 3、5、6:
单棵树表决就是它自己。多棵树的档位无非把这个过程重复 T 次再计票,复杂度全在读入不串位和平票规则上。
树模型推理在面试里对应两个高频问题。「随机森林为什么有效」:每棵树看到的数据与特征子集不同、各自犯不同的错,独立投票把方差平均掉——实现里那个多数表决循环就是这句话的代码形态。「决策树推理为什么快」:一个样本只走一条根到叶的路径,代价是树深 O(h),与训练集大小无关。能顺手补一句「工程上走树用循环不用递归,5000 节点的链状树会打爆默认递归栈」,这个细节比背十条八股更能证明真写过。
先用样例 1 的「5 ≤ 5 走左」验证比较方向;再造一棵两个类别票数相同的小树验证平票取小;RE 而非 WA 时,几乎可以直接断定是递归走树在深树用例上爆栈,改循环立刻解决。读入串位的症状很典型:前几个样本对、后面全错——检查每棵树是否先读了自己的节点数。
换个角度看,这道题是「把模型当数据」的第一课:树不是写死的 if/else 代码,而是一份被解释执行的节点表——模型部署引擎做的正是同一件事,只是格式从五元组换成了 ONNX 或 protobuf。类别值可以到十万量级且不连续,计票容器用哈希而非定长数组,也是同一个「模型是数据」的提醒。
# 多树推理:循环走树(深树递归会栈溢出),小于等于阈值走左;
# 多棵树表决取票数最多,平票取数值更小的类别,键=(−票数,类别)。
import sys
from collections import defaultdict
def solve() -> None:
data = sys.stdin.buffer.read().split()
if not data:
return
at = 0
t = int(data[at]); at += 1
d = int(data[at]); at += 1
q = int(data[at]); at += 1
trees = []
for _ in range(t):
s = int(data[at]); at += 1
nodes = [None] * (s + 1)
for i in range(1, s + 1):
typ = data[at].decode(); at += 1
if typ == "L":
nodes[i] = ("L", int(data[at])); at += 1
else:
feature = int(data[at]) - 1; at += 1
threshold = float(data[at]); at += 1
left = int(data[at]); at += 1
right = int(data[at]); at += 1
nodes[i] = ("N", feature, threshold, left, right)
trees.append(nodes)
ans = []
for _ in range(q):
x = [float(data[at + i]) for i in range(d)]
at += d
votes = defaultdict(int)
for nodes in trees:
cur = 1
while nodes[cur][0] != "L":
_, feature, threshold, left, right = nodes[cur]
cur = left if x[feature] <= threshold else right
votes[nodes[cur][1]] += 1
ans.append(str(min(votes, key=lambda label: (-votes[label], label))))
print("\n".join(ans))
if __name__ == "__main__":
solve()
登录后可查看你在本题的历史提交,以及每次的各用例通过情况。