如何聚类PyTorch道路车道检测的分散预测结果?
车道检测预测点聚类方案(解决分散点问题)
你提供的预测结果图显示大量分散的白色点分布在道路区域,目标效果则是只保留了集中的几条车道线点,分散点被完全移除。结合你的场景,以下是针对性的解决方案:
先明确数据格式
你的预测张量形状是[1, 1, 80, 120],应该是单batch、单通道的预测热力图:80对应图像的y轴高度(行),120对应x轴宽度(列),每个元素代表该位置属于车道的置信度。第一步必须先把高置信度的有效点提取出来,不然聚类会被大量低置信的噪声点干扰。
为什么KNN没用?
KNN是分类/回归算法,并非聚类算法——你可能用它尝试做异常点检测,但它对这种密度不均的车道点场景适配性差。推荐用DBSCAN(基于密度的空间聚类),它不需要预先指定聚类数量,还能自动识别并过滤分散的噪声点,完美匹配你的需求。
具体实现步骤
1. 提取有效候选点
先把PyTorch张量转换成二维点集:
import torch import numpy as np # 假设pred是你的预测张量,形状[1,1,80,120] pred = torch.randn(1,1,80,120) # 示例数据 pred = torch.sigmoid(pred) # 二分类预测需转成置信度 # 提取置信度高于阈值的点(阈值根据模型调整,比如0.6) conf_threshold = 0.6 pred_np = pred.squeeze().cpu().numpy() # 去掉前两个维度,得到(80,120) y_coords, x_coords = np.where(pred_np > conf_threshold) points = np.stack([x_coords, y_coords], axis=1) # 转成(N,2)的点集,格式为(x,y)
2. 用DBSCAN聚类过滤噪声
用scikit-learn的DBSCAN实现,核心是调整eps(邻域半径)和min_samples(邻域内最少点数)两个参数:
from sklearn.cluster import DBSCAN # 初始化DBSCAN:eps控制聚类紧密程度,min_samples控制簇的最小规模 # 需根据点分布微调,先试eps=5,min_samples=3 dbscan = DBSCAN(eps=5, min_samples=3) labels = dbscan.fit_predict(points) # 筛选出属于簇的点(label != -1代表不是噪声) valid_points = points[labels != -1]
3. 还原成热力图输出(可选)
如果需要把筛选后的点还原成和原预测同形状的张量:
clean_pred = np.zeros_like(pred_np) for x, y in valid_points: clean_pred[y, x] = pred_np[y, x] # 保留原置信度 # 转回PyTorch张量 clean_pred_tensor = torch.tensor(clean_pred).unsqueeze(0).unsqueeze(0) # 回到[1,1,80,120]
进阶优化:结合车道先验
车道线是连续的曲线/直线,同一车道的点在y方向应该连续分布。可以在聚类后,对每个簇做曲线拟合,进一步过滤偏离车道线的点:
from sklearn.preprocessing import PolynomialFeatures from sklearn.linear_model import LinearRegression # 遍历每个簇 unique_labels = np.unique(labels[labels != -1]) final_points = [] for label in unique_labels: cluster_points = points[labels == label] # 用二次多项式拟合车道线(y作为自变量,x作为因变量) X = cluster_points[:,1].reshape(-1,1) # y坐标 y = cluster_points[:,0] # x坐标 poly = PolynomialFeatures(degree=2) X_poly = poly.fit_transform(X) model = LinearRegression() model.fit(X_poly, y) # 预测每个y对应的x,计算误差,保留误差小的点 pred_x = model.predict(X_poly) error = np.abs(pred_x - y) final_points.extend(cluster_points[error < 2]) # 误差阈值2,可调整
效果说明
通过以上步骤,就能实现你想要的效果:只保留密集的车道点,过滤掉分散的噪声点。DBSCAN的参数需要根据数据集微调:如果还有分散点,就减小eps或者增大min_samples;如果车道点被误删,就增大eps或者减小min_samples。
内容的提问来源于stack exchange,提问作者Mustafa Uğur Baskın
相关产品推荐
相关产品推荐

