Python实现向量量化码本时出现float不可迭代错误如何解决?
问题原因
float类型不可迭代报错由变量类型覆盖导致:
- 你最初将
eu_dist15定义为存储所有距离的空列表 - 后续遍历
new_set计算距离时,直接使用eu_dist15= euclidean_distance(row15, i)赋值,把列表变量覆盖为了单个浮点型的距离值 - 最终尝试遍历
eu_dist15这个单值float时,就触发了不可迭代的报错
另外你的欧氏距离计算函数存在逻辑缺陷:for i in range(len(row1)-1)会漏掉向量最后一个维度的差值计算,4维向量计算欧氏距离需要遍历所有4个维度,需要去掉-1。
修复方案
将计算距离的循环逻辑修改为向列表追加元素,不要直接覆盖变量:
# 先修正欧氏距离函数 def euclidean_distance(row1, row2): distance = 0.0 for i in range(len(row1)): # 去掉-1,计算所有维度 distance += (row1[i] - row2[i])**2 return math.sqrt(distance) # 修正距离计算逻辑 for i in new_set: current_dist = euclidean_distance(row15, i) eu_dist15.append(current_dist) # 追加到列表,不覆盖原变量 print(current_dist) # 此时eu_dist15是列表,可正常遍历 for i in eu_dist15: print(i)
优化建议
如果你需要计算8个初始码本和所有样本的距离,不需要单独定义8个距离列表,可以用numpy批量运算简化代码:
import numpy as np # 转成numpy数组直接批量计算欧氏距离 codebook_arr = np.array(codebook_init) new_set_arr = np.array(new_set) # 得到形状为(8, 16)的距离矩阵,每行对应一个码本和所有样本的距离 dist_matrix = np.linalg.norm(codebook_arr[:, None] - new_set_arr, axis=2)
内容的提问来源于stack exchange,提问作者Random Thots
相关产品推荐
相关产品推荐

