Pandas筛选CSV时序数据适配TensorFlow TimeSeriesGenerator方法
Pandas + NumPy 实现时序数据过滤方案
整体实现优先保证从根源规避跨观测链取数、true标签出现在序列中间两类非法问题,全程不需要依赖复杂的时序窗口函数,逻辑可直接调试修改。
第一步:读取原始数据、拆分独立观测链
b r e a k是观测链的硬分隔符,绝对不能跨分隔符取数,因此先逐行读文件拆分独立链,避免后续滑窗越界:
import pandas as pd import numpy as np SEQ_LEN = 3 # 自定义固定序列长度 MIN_PRE_FALSE = SEQ_LEN - 1 # 序列末尾标签前最少需要的false数量,默认是长度-1 chains = [] current_chain_cache = [] # 逐行读取原始CSV with open("your_dataset.csv", "r", encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue # 遇到分隔符就封存当前链,重置缓存 if line == "b r e a k": if current_chain_cache: chains.append(current_chain_cache) current_chain_cache = [] continue # 解析普通数据行,统一转成数值类型 obs1, obs2, tag = line.split(",") current_chain_cache.append({ "obs1": float(obs1), "obs2": float(obs2), "tag": 1 if tag.strip().lower() == "true" else 0 }) # 处理文件末尾未追加break的最后一条链 if current_chain_cache: chains.append(current_chain_cache)
第二步:逐链提取合法固定长度序列
校验逻辑完全对齐要求:
- 链长度不足固定序列长度的直接丢弃
- 若观测链长度超过固定值,优先取链尾部的片段做校验(可按需修改为全滑窗提取)
- 强制校验序列前
MIN_PRE_FALSE位的标签全为false,从规则上保证true标签只能出现在序列最后一位
valid_seq = [] for chain in chains: chain_len = len(chain) # 凑不够固定长度直接跳过 if chain_len < SEQ_LEN: continue # --- 仅取链尾部符合长度的片段,符合需求里“丢弃头部多余数据”的要求 --- candidate = chain[-SEQ_LEN:] # 校验前置位无true标签 pre_tag_sum = sum([item["tag"] for item in candidate[:-1]]) if pre_tag_sum == 0: # 转成模型输入格式:特征是[SEQ_LEN, 2]的数组,标签是最后一位的tag seq_feature = np.array([[item["obs1"], item["obs2"]] for item in candidate]) seq_label = candidate[-1]["tag"] valid_seq.append((seq_feature, seq_label)) # --- 如果需要提取单链内所有合法片段,替换上面的代码为下面的滑窗逻辑即可 --- # for start in range(chain_len - SEQ_LEN + 1): # candidate = chain[start:start+SEQ_LEN] # pre_tag_sum = sum([item["tag"] for item in candidate[:-1]]) # if pre_tag_sum == 0: # seq_feature = np.array([[item["obs1"], item["obs2"]] for item in candidate]) # seq_label = candidate[-1]["tag"] # valid_seq.append((seq_feature, seq_label))
第三步:输出适配TensorFlow输入格式
把校验通过的序列转成NumPy数组,可直接喂给模型,不需要额外依赖TimeSeriesGenerator(已经提前做好窗口切分,从根源避免非法序列):
# 特征数组shape:(合法样本数, 序列长度, 观测特征数) X = np.stack([item[0] for item in valid_seq]) # 标签数组shape:(合法样本数,) y = np.array([item[1] for item in valid_seq])
注意事项:不要直接把全量CSV读成DataFrame后调用
TimeSeriesGenerator默认滑窗逻辑,该方法无法识别b r e a k分隔符,一定会生成跨观测链的非法序列,先拆链再校验的逻辑可以100%规避这个问题。如果需要调整前置false的最少数量,直接修改MIN_PRE_FALSE参数、对应调整校验时的切片范围即可。
内容的提问来源于stack exchange,提问作者studenprogrammer
相关产品推荐
相关产品推荐

