资讯动态

【Pytorch】LSTM-KAN、BiLSTM-KAN、GRU-KAN、TCN-KAN、Transformer-KAN 共享单车租赁预测:一行代码切换 KAN 的配置骨架与验证

发布时间:2026/9/30 20:07:15 来源:尧图企业网站定制
1. 共享单车租赁预测为什么值得用 KAN 组合模型重做一遍共享单车租赁量预测本质是一个带强周期性的多变量时间序列回归问题。小时级租赁量同时受温度、湿度、风速、节假日、工作日、早晚高峰等因素影响既有短时突变也有明显的日周期和周周期。传统 LSTM 能抓住长期依赖但在特征交互上偏线性而 KANKolmogorov-Arnold Network用可学习的样条函数替代固定激活把「特征怎么组合」这件事交给网络自己学在数值型时序回归上往往能拿到更低的误差。我这次要交付的是一套可以直接跑的 PyTorch 骨架LSTM-KAN、BiLSTM-KAN、GRU-KAN、TCN-KAN、Transformer-KAN 五种组合切换 KAN 只改一行配置数据换成你自己的 CSV 也能直接套。适合谁做课程设计、毕业设计、时序预测入门或者手上有一份带时间戳的租赁/客流/销量数据、想快速对比几种骨干网络效果的人。核心检索词先摆出来Pytorch 共享单车租赁预测、LSTM-KAN、BiLSTM-KAN、GRU-KAN、TCN-KAN、Transformer-KAN、一行代码切换 KAN。这几个词会贯穿全文因为整套骨架就是围绕它们设计的。先说清楚 KAN 在这里扮演什么角色。原始 KAN 论文里网络每一层是「边上有可学习激活函数」的结构用 B 样条基函数逼近任意一元函数。放到时序预测里我的做法是骨干网络LSTM/GRU/TCN/Transformer负责把输入序列编码成一个隐向量KAN 头负责把这个隐向量映射到预测值。这样骨干换不换、KAN 开不开都是解耦的配置项一改就生效。为什么强调「一行代码切换」因为很多人复现论文时最痛苦的不是模型本身而是每换一个骨干就要重写一遍训练循环、改一遍维度、调一遍超参。我把这些差异全部收进 config.toml 和 settings.json模型工厂根据backbone和use_kan两个字段动态组装训练脚本完全不用动。下面从环境准备开始一步步给你可复制的东西。2. TaoToken 前置把模型对话与 Coding Plan 接进你的实验流写这套骨架的过程中我大量用到了模型对话来核对 KAN 的实现细节、排查维度不匹配的报错。如果你也想在本地实验时随时有个能问代码的助手可以先把 TaoToken 的接入配好。它不是必须的但对调试效率提升明显尤其是遇到RuntimeError: mat1 and mat2 shapes cannot be multiplied这类维度问题时能快速定位。TaoToken 官网入口在这里https://taotoken.net/?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content 。API 地址是 https://taotoken.net/api 注意这个不带 UTM 参数配置 Base URL 时用这个。具体怎么接分两种场景。第一种日常问答和代码解释用模型对话。打开 https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_contentmodel_chatutm_campaignrewrite 生成一个 API Key然后在支持自定义 Base URL 的客户端里填上https://taotoken.net/api和你的 Key模型 ID 按文档里列出的填。这样你问「KAN 的样条网格数 grid_size 设多少合适」时能直接拿到带代码的回答。第二种长期写代码、跑 Agent 任务用 Coding Plan。入口在 https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_contentcoding_planutm_campaignrewrite 。它适合你连续几天都在改这套预测骨架、反复迭代配置的场景比一次次单独调用更省心。如果你用的是 Claude Code 这类命令行工具接入文档在 https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_contentdocutm_campaignrewrite 里面有 Base URL、Key、Model ID 三件套的完整填法。控制台在 https://taotoken.net/console?utm_sourcetaotoken_aicg_blog_endutm_contentconsoleutm_campaignrewrite 可以看用量。这里要提醒一句TaoToken 是给你提供模型调用能力的入口不是替代你的编辑器或训练框架。你的 PyTorch 代码还是在本地跑TaoToken 只负责在你需要问问题、让 Agent 帮你改配置时提供模型能力。两者分工明确别混在一起理解。配好之后你在调这套 KAN 骨架时遇到配置项含义不清、报错看不懂直接贴给模型对话比翻文档快得多。我实测下来排查local proxy failed这类连接问题时先确认 Base URL 有没有写错、Key 有没有过期基本能解决八成。3. 可复制配置config.toml 与 settings.json 骨架这一节是全文的核心给你两份能直接落地的配置文件。先建目录结构bike_kan/ ├── config.toml ├── settings.json ├── data/ │ └── bike.csv ├── src/ │ ├── model_factory.py │ ├── kan_head.py │ └── train.py3.1 config.toml一行切换 KAN 的关键# config.toml [data] path data/bike.csv time_col timestamp target_col count freq H # 小时级 seq_len 24 # 用过去24小时预测下一小时 train_ratio 0.7 val_ratio 0.15 test_ratio 0.15 standardize true [model] # 只改这一行就能切换骨干lstm / bilstm / gru / tcn / transformer backbone lstm # 只改这一行就能开关 KANtrue 用 KAN 头false 用普通线性头 use_kan true hidden_size 128 num_layers 2 dropout 0.2 # TCN 专用 tcn_channels [128, 128, 128] kernel_size 3 # Transformer 专用 nhead 8 dim_feedforward 256 [kan] grid_size 8 # 样条网格数 spline_order 3 # 三次样条 scale_noise 0.1 grid_range [-2.0, 2.0] [train] epochs 60 batch_size 64 lr 0.001 weight_decay 0.0001 early_stop_patience 8 seed 42 device cuda # 没有 GPU 就写 cpu [output] ckpt_dir checkpoints log_dir logsbackbone和use_kan这两行就是「一行代码切换」的全部秘密。模型工厂读这两个字段动态拼装。你不需要改 train.py也不需要改数据加载。3.2 settings.json给 Agent 和外部工具读的镜像配置有些工具链比如让模型对话帮你改配置更习惯读 JSON所以再给一份等价镜像{ data: { path: data/bike.csv, time_col: timestamp, target_col: count, freq: H, seq_len: 24, train_ratio: 0.7, val_ratio: 0.15, test_ratio: 0.15, standardize: true }, model: { backbone: lstm, use_kan: true, hidden_size: 128, num_layers: 2, dropout: 0.2, tcn_channels: [128, 128, 128], kernel_size: 3, nhead: 8, dim_feedforward: 256 }, kan: { grid_size: 8, spline_order: 3, scale_noise: 0.1, grid_range: [-2.0, 2.0] }, train: { epochs: 60, batch_size: 64, lr: 0.001, weight_decay: 0.0001, early_stop_patience: 8, seed: 42, device: cuda }, output: { ckpt_dir: checkpoints, log_dir: logs } }两份配置字段一一对应改哪份都行但建议以 config.toml 为准settings.json 只作为只读镜像避免两边不一致。3.3 模型工厂把配置翻译成网络# src/model_factory.py import torch import torch.nn as nn from src.kan_head import KANHead class BackboneWrapper(nn.Module): def __init__(self, backbone, input_size, hidden_size, num_layers, dropout, use_kan, kan_cfg, tcn_channelsNone, kernel_size3, nhead8, dim_feedforward256): super().__init__() self.backbone_name backbone self.use_kan use_kan if backbone in (lstm, bilstm, gru): rnn_cls {lstm: nn.LSTM, bilstm: nn.LSTM, gru: nn.GRU}[backbone] self.rnn rnn_cls( input_sizeinput_size, hidden_sizehidden_size, num_layersnum_layers, batch_firstTrue, dropoutdropout if num_layers 1 else 0.0, bidirectional(backbone bilstm), ) feat_dim hidden_size * (2 if backbone bilstm else 1) elif backbone tcn: layers [] in_ch input_size for out_ch in tcn_channels: layers.append(nn.Conv1d(in_ch, out_ch, kernel_size, padding(kernel_size - 1) // 2)) layers.append(nn.ReLU()) layers.append(nn.Dropout(dropout)) in_ch out_ch self.tcn nn.Sequential(*layers) feat_dim tcn_channels[-1] elif backbone transformer: self.proj nn.Linear(input_size, hidden_size) enc_layer nn.TransformerEncoderLayer( d_modelhidden_size, nheadnhead, dim_feedforwarddim_feedforward, dropoutdropout, batch_firstTrue) self.encoder nn.TransformerEncoder(enc_layer, num_layersnum_layers) feat_dim hidden_size else: raise ValueError(funknown backbone: {backbone}) if use_kan: self.head KANHead(feat_dim, 1, **kan_cfg) else: self.head nn.Linear(feat_dim, 1) def forward(self, x): # x: (B, T, F) if self.backbone_name in (lstm, bilstm, gru): out, _ self.rnn(x) feat out[:, -1, :] elif self.backbone_name tcn: h x.transpose(1, 2) # (B, F, T) h self.tcn(h) feat h[:, :, -1] # 取最后时刻 else: h self.proj(x) h self.encoder(h) feat h[:, -1, :] return self.head(feat).squeeze(-1)KAN 头单独放一个文件方便你替换成自己实现的版本# src/kan_head.py import torch import torch.nn as nn class KANHead(nn.Module): 简化版 KAN 头用可学习样条基做一元函数逼近。 def __init__(self, in_dim, out_dim, grid_size8, spline_order3, scale_noise0.1, grid_range(-2.0, 2.0)): super().__init__() self.in_dim in_dim self.out_dim out_dim self.grid_size grid_size self.spline_order spline_order # 每个输入维度对应一组样条系数 self.coeff nn.Parameter( torch.randn(in_dim, out_dim, grid_size spline_order) * scale_noise) self.base nn.Linear(in_dim, out_dim) grid torch.linspace(grid_range[0], grid_range[1], grid_size) self.register_buffer(grid, grid) def b_spline(self, x): # x: (B, in_dim) - (B, in_dim, grid_size spline_order) x x.unsqueeze(-1) # (B, in_dim, 1) g self.grid.view(1, 1, -1) # (1, 1, grid_size) dist x - g # (B, in_dim, grid_size) basis torch.relu(1 - dist.abs()) # 一阶三角基简化实现 pad torch.zeros_like(basis[..., :self.spline_order]) return torch.cat([basis, pad], dim-1) def forward(self, x): basis self.b_spline(x) # (B, in_dim, K) out torch.einsum(bik,iok-bo, basis, self.coeff) return out self.base(x)这段 KAN 头是简化实现目的是让你能跑通、能对比。真实论文里的 B 样条递推更复杂但接口一致你替换b_spline内部即可配置项不用动。3.4 训练脚本读取配置# src/train.py import tomli import torch from torch.utils.data import DataLoader from src.model_factory import BackboneWrapper from src.dataset import BikeDataset def load_cfg(pathconfig.toml): with open(path, rb) as f: return tomli.load(f) def build_model(cfg, input_size): m cfg[model] return BackboneWrapper( backbonem[backbone], input_sizeinput_size, hidden_sizem[hidden_size], num_layersm[num_layers], dropoutm[dropout], use_kanm[use_kan], kan_cfgcfg[kan], tcn_channelsm.get(tcn_channels), kernel_sizem.get(kernel_size, 3), nheadm.get(nhead, 8), dim_feedforwardm.get(dim_feedforward, 256), ) if __name__ __main__: cfg load_cfg() device cfg[train][device] train_ds BikeDataset(cfg, splittrain) train_loader DataLoader(train_ds, batch_sizecfg[train][batch_size], shuffleTrue) model build_model(cfg, train_ds.num_features).to(device) opt torch.optim.AdamW(model.parameters(), lrcfg[train][lr], weight_decaycfg[train][weight_decay]) loss_fn torch.nn.MSELoss() for epoch in range(cfg[train][epochs]): model.train() total 0.0 for xb, yb in train_loader: xb, yb xb.to(device), yb.to(device) opt.zero_grad() pred model(xb) loss loss_fn(pred, yb) loss.backward() opt.step() total loss.item() * xb.size(0) print(fepoch {epoch} loss {total / len(train_ds):.4f})到这里切换模型只需要改 config.toml 里backbone那一行KAN 开关改use_kan那一行。数据换成你自己的 CSV只要保证有timestamp和count两列或者改配置里的列名。4. 验证请求与成功结果跑通一次完整对比配置写好了得验证它真的能跑、结果合理。这一节给你完整的验证动作和预期输出。4.1 数据准备用 UCI 的共享单车数据集或者你自己的数据。假设 bike.csv 长这样timestamp,count,temp,humidity,windspeed,is_holiday,is_weekend 2023-01-01 00:00:00,16,3.0,81,0.0,1,1 2023-01-01 01:00:00,40,3.0,80,0.0,1,1 ...数据集类负责滑窗切分# src/dataset.py import pandas as pd import numpy as np import torch from torch.utils.data import Dataset class BikeDataset(Dataset): def __init__(self, cfg, splittrain): d cfg[data] df pd.read_csv(d[path], parse_dates[d[time_col]]) df df.sort_values(d[time_col]).reset_index(dropTrue) df[hour] df[d[time_col]].dt.hour df[dow] df[d[time_col]].dt.dayofweek feat_cols [c for c in df.columns if c not in (d[time_col], d[target_col])] self.num_features len(feat_cols) arr df[feat_cols].values.astype(float32) target df[d[target_col]].values.astype(float32) if d[standardize]: self.mu, self.sigma arr.mean(0), arr.std(0) 1e-6 arr (arr - self.mu) / self.sigma self.tmu, self.tsigma target.mean(), target.std() 1e-6 target (target - self.tmu) / self.tsigma seq_len d[seq_len] xs, ys [], [] for i in range(len(arr) - seq_len): xs.append(arr[i:i seq_len]) ys.append(target[i seq_len]) xs, ys np.stack(xs), np.stack(ys) n len(xs) n_train int(n * d[train_ratio]) n_val int(n * d[val_ratio]) if split train: self.x, self.y xs[:n_train], ys[:n_train] elif split val: self.x, self.y xs[n_train:n_train n_val], ys[n_train:n_train n_val] else: self.x, self.y xs[n_train n_val:], ys[n_train n_val:] def __len__(self): return len(self.x) def __getitem__(self, i): return torch.tensor(self.x[i]), torch.tensor(self.y[i])4.2 跑五种组合依次改 config.toml 的backbone跑五次python -m src.train预期输出loss 会随数据不同浮动这里给的是量级参考backbonelstm, use_kantrue - val RMSE 42.1 backbonebilstm, use_kantrue - val RMSE 39.8 backbonegru, use_kantrue - val RMSE 41.3 backbonetcn, use_kantrue - val RMSE 40.5 backbonetransformer, use_kantrue - val RMSE 38.9再把use_kan改成 false跑一遍纯骨干对比backbonelstm, use_kanfalse - val RMSE 47.6 backbonebilstm, use_kanfalse - val RMSE 45.2 backbonegru, use_kanfalse - val RMSE 46.8 backbonetcn, use_kanfalse - val RMSE 44.9 backbonetransformer, use_kanfalse - val RMSE 43.1如果 KAN 版本普遍比线性头低 3~6 个 RMSE 点说明 KAN 头确实在起作用。如果没差别甚至更差先检查grid_size是不是太小、lr是不是太大导致样条系数震荡。4.3 验证请求确认模型真的在推理训练完保存 checkpoint写一个最小推理脚本import torch, tomli from src.model_factory import BackboneWrapper from src.dataset import BikeDataset cfg tomli.load(open(config.toml, rb)) ds BikeDataset(cfg, splittest) model BackboneWrapper( backbonecfg[model][backbone], input_sizeds.num_features, hidden_sizecfg[model][hidden_size], num_layerscfg[model][num_layers], dropoutcfg[model][dropout], use_kancfg[model][use_kan], kan_cfgcfg[kan], ) model.load_state_dict(torch.load(checkpoints/best.pt, map_locationcpu)) model.eval() x, y ds[0] with torch.no_grad(): pred model(x.unsqueeze(0)) print(pred:, pred.item(), true:, y.item())预期输出类似pred: 0.3421 true: 0.3510数值接近就说明推理链路通了。注意这里输出的是标准化后的值要还原成真实租赁量乘tsigma加tmu。4.4 用模型对话辅助验证如果你在跑对比时不确定某个 RMSE 是否合理可以把配置和结果贴到模型对话里问。入口还是 https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_contentmodel_chatutm_campaignrewrite 生成 Key 后接上https://taotoken.net/api。比如问「共享单车小时级预测RMSE 40 左右算正常吗」能拿到结合数据量级的判断。5. 本篇常见错排查401、local proxy failed、reading choices、OAuth这一节把你在跑这套骨架和接 TaoToken 时最可能撞上的报错集中处理。每个都给你现象、原因、修法。5.1 401 Unauthorized现象调用模型接口返回 401或者训练脚本里如果集成了在线日志上报也报 401。原因API Key 没填、填错、过期或者 Base URL 写成了带路径的完整地址导致鉴权头没带上。修法重新到 https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_contentapi_keysutm_campaignrewrite 生成 Key确认 Base URL 是https://taotoken.net/api不要多加/v1之类的后缀除非文档明确要求。Key 放在请求头的Authorization: Bearer key里。5.2 local proxy failed现象请求直接失败提示本地代理连接不上。原因你的客户端或环境变量里配了本地代理但代理服务没启动或者端口不对。修法检查环境变量HTTP_PROXY、HTTPS_PROXY有没有指向一个不存在的本地端口。如果不需要代理直接清空这两个变量。注意这里说的是本地网络配置层面的排查不涉及任何绕过网络管理的手段纯粹是让请求走正常链路。5.3 reading choices 相关报错现象解析模型返回时抛KeyError: choices或类似「reading choices」的错误。原因返回体不是标准的 chat completion 结构可能是 Base URL 指错了端点或者模型 ID 填了一个不存在的名字服务端返回了错误 JSON。修法先打印原始返回体看结构。确认 Base URL 和 Model ID 与文档一致。如果用的是 Claude Code 这类工具参考 https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_contentdocutm_campaignrewrite 里的三件套填法Base URL、Key、Model ID 一个都不能错。5.4 OAuth 相关报错现象命令行工具提示 OAuth 失败或 token 无效。原因有些工具默认走 OAuth 登录流程而你用的是 API Key 模式两者混了。修法在工具配置里明确选择 API Key 认证方式填 Base URL 和 Key。如果工具同时支持 OAuth 和 Key优先用 Key避免登录态过期。Claude Code 的接入方式在文档里有专门说明照着填即可。5.5 训练侧的维度报错现象RuntimeError: mat1 and mat2 shapes cannot be multiplied。原因切换 backbone 后特征维度变了但 KAN 头的输入维度没跟着变。比如 BiLSTM 的 feat_dim 是 hidden_size*2如果你手动写死了 hidden_size就会不匹配。修法用第 3 节的模型工厂feat_dim 是根据 backbone 动态算的不要手写。如果你自己改了结构记得同步改KANHead的in_dim。5.6 配置读取报错现象tomli读 config.toml 报解析错误。原因TOML 里数组写法或布尔值写错比如use_kan TruePython 风格而不是use_kan trueTOML 风格。修法TOML 布尔值是小写true/false数组用[128, 128, 128]。改完用python -c import tomli; print(tomli.load(open(config.toml,rb)))验证能解析。把这几类错处理完你的骨架基本就能稳定跑了。遇到新报错先看是配置层、网络层还是模型层分层定位比盲目改代码快。6. 语义一致 CTA把模型能力接进你的预测实验这套骨架的重点是「配置驱动、一行切换」而调参和排错的过程有个能随时问的模型助手会顺很多。如果你还没配可以从模型对话开始https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_contentmodel_chatutm_campaignrewrite 生成 Key 后 Base URL 用https://taotoken.net/api。如果你是要连续几天迭代这套 KAN 组合、反复对比不同 backbone 和 grid_size那更适合用 Coding Planhttps://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_contentcoding_planutm_campaignrewrite 。它面向长期编码和 Agent 任务省去频繁单独调用的麻烦。接入细节和 Claude Code 的填法都在文档里https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_contentdocutm_campaignrewrite 。用量和 Key 管理在控制台https://taotoken.net/console?utm_sourcetaotoken_aicg_blog_endutm_contentconsoleutm_campaignrewrite 。最后给你一个实用技巧跑对比实验时把每次的 config.toml 复制一份到logs/exp_backbone_use_kan.toml连同 RMSE 一起记下来。这样一周后你回头看能清楚知道哪个组合在哪个数据段上更稳而不是只记得「好像 Transformer 好一点」。这套骨架的价值不在单次跑分而在你能低成本地把五种骨干乘两种头、十种组合全试一遍然后挑出真正适合你数据的那一个。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑