训练Estimating People Flows数据集时索引转int类型错误的解决咨询
问题:数据集训练时的索引类型错误
在训练《Estimating People Flows to Better Count Them in Crowded Scenes》论文的数据集时,执行index = int(img_name.split('.')[0])代码时触发类型错误——因为部分图片名称格式为130(1)、049(1),无法直接转换为整数,报错提示:
ValueError: invalid literal for int() with base 10: '49(1)'
多次运行时,报错的图片名称会变化,比如首次出现049(1),后续出现130(1)。
详细报错栈
Traceback (most recent call last): File "train.py", line 253, in <module> main() File "train.py", line 58, in main train(train_list, model, criterion, optimizer, epoch) File "train.py", line 91, in train for i,(prev_img, img, post_img, prev_target, target, post_target ) in enumerate(train_loader): File "/root/miniconda3/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 435, in __next__ data = self._next_data() File "/root/miniconda3/lib/python3.8/site-packages/torch/utils/data/dataloader.py", line 475, in _next_data data = self._dataset_fetcher.fetch(index) # may raise StopIteration File "/root/miniconda3/lib/python3.8/site-packages/torch/utils/data/_utils/fetch.py", line 44, in fetch data = [self.dataset[idx] for idx in possibly_batched_index] File "/root/miniconda3/lib/python3.8/site-packages/torch/utils/data/_utils/fetch.py", line 44, in <listcomp> data = [self.dataset[idx] for idx in possibly_batched_index] File "/root/autodl-tmp/project/dataset.py", line 30, in __getitem__ prev_img,img,post_img,prev_target, target, post_target = load_data(img_path,self.train) File "/root/autodl-tmp/project/image.py", line 12, in load_data index = int(img_name.split('.')[0]) ValueError: invalid literal for int() with base 10: '49(1)'
原始代码片段
def load_data(img_path,train = True): img_folder = os.path.dirname(img_path) img_name = os.path.basename(img_path) # print(img_name.split('.')[0]) index = int(img_name.split('.')[0]) # 出错行
用户尝试与疑问
最初尝试将(和)替换为数字,但发现会导致索引与图片标签不匹配,无法找到对应图片。疑问:是否只能修改数据集中的图片标签来解决问题?
解决方案
不需要修改数据集的图片名称,只需修改解析索引的逻辑,提取出括号前的核心数字即可,以下是几种可行方法:
方法1:按括号分割提取数字
直接通过字符串分割,取括号前的部分转换为整数:
def load_data(img_path,train = True): img_folder = os.path.dirname(img_path) img_name = os.path.basename(img_path) name_part = img_name.split('.')[0] # 按(分割,取第一部分作为索引源 index = int(name_part.split('(')[0])
方法2:正则表达式提取连续数字
如果存在更复杂的命名格式,用正则提取开头的连续数字:
import re def load_data(img_path,train = True): img_folder = os.path.dirname(img_path) img_name = os.path.basename(img_path) name_part = img_name.split('.')[0] # 匹配开头的所有连续数字 num_match = re.match(r'^\d+', name_part) if num_match: index = int(num_match.group()) else: # 处理无有效数字的异常情况 raise ValueError(f"无法从{name_part}中提取索引数字")
方法3:移除括号及内部内容
如果括号内是重复标记类内容,直接移除括号和内部内容后转整数:
import re def load_data(img_path,train = True): img_folder = os.path.dirname(img_path) img_name = os.path.basename(img_path) name_part = img_name.split('.')[0] # 移除所有()及其中的内容 clean_name = re.sub(r'\(.*?\)', '', name_part) index = int(clean_name)
以上方法均无需修改数据集文件名称,仅在代码层面处理即可,既保证索引与图片名称的对应关系,又能获取到正确的整数索引。
内容的提问来源于stack exchange,提问作者rayla
相关产品推荐
相关产品推荐

