Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
46 commits
Select commit Hold shift + click to select a range
20d6cf2
add paddle parallel
DrownFish19 Oct 16, 2023
790e398
fix corrstn train
DrownFish19 Oct 16, 2023
54c3a98
add test before finetune
DrownFish19 Oct 16, 2023
11066ed
update interpolate
DrownFish19 Oct 16, 2023
d07b90d
set default model
DrownFish19 Oct 16, 2023
1d41e7d
Merge remote-tracking branch 'refs/remotes/origin/dev-corrstn_ddeint'…
DrownFish19 Oct 16, 2023
d10d51f
fix interpolation
DrownFish19 Oct 16, 2023
07f042d
update interplate
DrownFish19 Oct 16, 2023
1a55704
fix
DrownFish19 Oct 17, 2023
bc04fa6
tmp upload
DrownFish19 Oct 17, 2023
50976ab
update test cases
DrownFish19 Oct 17, 2023
b4f6d97
update demo
DrownFish19 Oct 17, 2023
7133b9b
update demo
DrownFish19 Oct 17, 2023
dca43df
temp update
DrownFish19 Oct 18, 2023
f2fc010
temp update
DrownFish19 Oct 24, 2023
53e4605
update
DrownFish19 Oct 24, 2023
3ffc65b
update model
DrownFish19 Oct 25, 2023
8e5744a
update history data length
DrownFish19 Oct 29, 2023
5f6f791
update requirement.txt
DrownFish19 Oct 29, 2023
4ce1361
update .gitignore
DrownFish19 Oct 29, 2023
02d49b9
combine ddeint into corrstn
DrownFish19 Oct 30, 2023
725d7c1
add paddleviz for corrstn
DrownFish19 Oct 30, 2023
f3abf0c
combine ddeint into train_dde.py
DrownFish19 Oct 30, 2023
68b9492
update train strategy
DrownFish19 Oct 30, 2023
b3ae2e5
【confirmed】update train strategy
DrownFish19 Nov 1, 2023
52a9474
update submodule
DrownFish19 Nov 3, 2023
60c30a5
update his data select
DrownFish19 Nov 3, 2023
5656049
pre-process dataset
DrownFish19 Nov 3, 2023
0d32b6b
update decoder init
DrownFish19 Nov 3, 2023
7356bb3
update train
DrownFish19 Nov 3, 2023
b4b40c1
args add no_adj
DrownFish19 Nov 5, 2023
2f5b611
update requirements
DrownFish19 Nov 6, 2023
3881d0e
update train
DrownFish19 Nov 6, 2023
8f6aaf4
update model
DrownFish19 Nov 7, 2023
60a229f
update train
DrownFish19 Nov 7, 2023
6257ecc
fix
DrownFish19 Nov 7, 2023
198a67f
fix decoder_idx
DrownFish19 Nov 7, 2023
f294a5e
update trainer
DrownFish19 Nov 10, 2023
d95a199
update logging output
DrownFish19 Nov 11, 2023
0c753c0
update model
DrownFish19 Nov 11, 2023
a0ec8ee
update visualdl writer
DrownFish19 Nov 11, 2023
4e381ca
update train
DrownFish19 Nov 11, 2023
a91a22d
update dataset
DrownFish19 Nov 12, 2023
8ba08c2
update model
DrownFish19 Nov 12, 2023
4e4749c
update train
DrownFish19 Nov 12, 2023
7c9e0cf
update model
DrownFish19 Nov 14, 2023
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -134,3 +134,6 @@ FETCH_HEAD

# vscode
# .vscode

# experiments
experiments/*
14 changes: 14 additions & 0 deletions .vscode/launch.json
Original file line number Diff line number Diff line change
Expand Up @@ -70,5 +70,19 @@
"PYTHONPATH": "${workspaceFolder}",
}
},
{
"name": "example_CorrSTN_ddeint",
"type": "python",
"request": "launch",
"program": "${workspaceFolder}/example/CorrSTN/train_dde.py",
"console": "integratedTerminal",
"justMyCode": false,
"args": [],
"env": {
"CUDA_VISIBLE_DEVICES": "0",
"PYTHONPATH": "${workspaceFolder}",
}
},

]
}
2 changes: 1 addition & 1 deletion example/CorrSTN/TrafficFlowData
18 changes: 11 additions & 7 deletions example/CorrSTN/args.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,35 +31,39 @@


# model config
parser.add_argument("--model_name", type=str, default="WindSTN", help="model name")
parser.add_argument("--model_name", type=str, default="DDECorrSTN", help="model name")
parser.add_argument("--his_len", type=int, default=288, help="history data length")
parser.add_argument("--tgt_len", type=int, default=12, help="tgt data length")
parser.add_argument("--encoder_input_size", type=int, default=1)
parser.add_argument("--decoder_input_size", type=int, default=1)
parser.add_argument("--decoder_output_size", type=int, default=1)
parser.add_argument("--encoder_num_layers", type=int, default=4)
parser.add_argument("--decoder_num_layers", type=int, default=4)
parser.add_argument("--d_model", type=int, default=64)
parser.add_argument("--d_model", type=int, default=64, help="d_proj+d_sect*2")
parser.add_argument("--d_proj", type=int, default=32)
parser.add_argument("--d_sect", type=int, default=16)
parser.add_argument("--attention", type=str, default="Corr", help="Corr,Vanilla")
parser.add_argument("--split_seq", type=bool, default=False, help="split q k v")
parser.add_argument("--head", type=int, default=8, help="head")
parser.add_argument("--kernel_size", type=int, default=3, help="kernel_size")
parser.add_argument("--top_k", type=int, default=5, help="top_k")
parser.add_argument("--smooth_layer_num", type=int, default=1)

parser.add_argument("--no_adj", type=bool, default=False, help="no adj")

# train config
parser.add_argument("--learning_rate", type=float, default=1e-3)
parser.add_argument("--weight_decay", type=float, default=0.01)
parser.add_argument("--weight_decay", type=float, default=0.0)
parser.add_argument("--start_epoch", type=int, default=0, help="start epoch")
parser.add_argument("--train_epochs", type=int, default=75, help="train epochs")
parser.add_argument("--train_epochs", type=int, default=100, help="train epochs")
parser.add_argument("--finetune_epochs", type=int, default=50, help="finetune epochs")
parser.add_argument("--batch_size", type=int, default=16, help="batch_size")
parser.add_argument("--patience", type=int, default=8, help="early stopping patience")
parser.add_argument("--patience", type=int, default=15, help="early stopping patience")
parser.add_argument("--loss", type=str, default="mse", help="loss function")
parser.add_argument("--dropout", type=float, default=0.0, help="dropout")
parser.add_argument("--continue_training", type=bool, default=False, help="")
parser.add_argument("--fp16", type=bool, default=False, help="")
parser.add_argument("--distribute", type=bool, default=False, help="")

args = parser.parse_args("")
# args = parser.parse_args("")
args = parser.parse_args()
os.environ["CUDA_VISIBLE_DEVICES"] = args.devices
4 changes: 4 additions & 0 deletions example/CorrSTN/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,13 +125,15 @@ def __init__(
args.d_model,
(1, args.kernel_size),
padding=(0, self.padding_1DConv),
bias_attr=True,
)
else:
self.query_conv = nn.Conv2D(
args.d_model,
args.d_model,
(1, args.kernel_size),
padding=(0, self.padding_causal),
bias_attr=True,
)

if key_conv_type == "1DConv":
Expand All @@ -140,13 +142,15 @@ def __init__(
args.d_model,
(1, args.kernel_size),
padding=(0, self.padding_1DConv),
bias_attr=True,
)
else:
self.key_conv = nn.Conv2D(
args.d_model,
args.d_model,
(1, args.kernel_size),
padding=(0, self.padding_causal),
bias_attr=True,
)

self.dropout = nn.Dropout(p=args.dropout)
Expand Down
150 changes: 0 additions & 150 deletions example/CorrSTN/corrstn.ipynb

This file was deleted.

189 changes: 154 additions & 35 deletions example/CorrSTN/corrstn.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,21 @@
from copy import deepcopy

import paddle.nn as nn
import paddle
from attention import MultiHeadAttentionAwareTemporalContext
from embedding import SpatialPositionalEmbedding, TemporalPositionalEmbedding
from embedding import TemporalSectionEmbedding

# SpatialPositionalEmbedding,; TemporalPositionalEmbedding,
from endecoder import Decoder, DecoderLayer, Encoder, EncoderLayer
from graphconv import GCN, SpatialAttentionGCN
from graphconv import SpatialAttentionGCN

# GCN
from paddle import autograd, nn

from paddlexde.interpolation.interpolate import (
BezierSpline,
CubicHermiteSpline,
LinearInterpolation,
)


class CorrSTN(nn.Layer):
Expand All @@ -14,23 +25,25 @@ def __init__(self, training_args, adj_matrix, sc_matrix):
self.training_args = training_args

self.encoder_dense = nn.Linear(
training_args.encoder_input_size, training_args.d_model
training_args.encoder_input_size, training_args.d_proj
)
self.decoder_dense = nn.Linear(
training_args.decoder_input_size, training_args.d_model
)

self.encode_temporal_position = TemporalPositionalEmbedding(
training_args, max_len=training_args.his_len
)
self.decode_temporal_position = TemporalPositionalEmbedding(
training_args, max_len=training_args.tgt_len
training_args.decoder_input_size, training_args.d_proj
)

self.encode_spatial_position = SpatialPositionalEmbedding(
training_args, GCN(training_args, adj_matrix, sc_matrix)
)
self.decode_spatial_position = deepcopy(self.encode_spatial_position)
# self.encode_temporal_position = TemporalPositionalEmbedding(
# training_args, max_len=training_args.his_len
# )
# self.decode_temporal_position = TemporalPositionalEmbedding(
# training_args, max_len=training_args.tgt_len
# )
self.temporal_section_week = TemporalSectionEmbedding(training_args, 7, axis=1)
self.temporal_section_day = TemporalSectionEmbedding(training_args, 288, axis=2)
# self.encode_spatial_position = SpatialPositionalEmbedding(
# training_args,
# GCN(training_args, training_args.d_proj, adj_matrix, sc_matrix),
# )
# self.decode_spatial_position = deepcopy(self.encode_spatial_position)

attn_ss = MultiHeadAttentionAwareTemporalContext(
args=training_args,
Expand Down Expand Up @@ -83,36 +96,103 @@ def __init__(self, training_args, adj_matrix, sc_matrix):
training_args.d_model, training_args.decoder_output_size
)

def encode(self, src, lookup_index):
src_dense = self.encoder_dense(src)
src_tp_embedding = self.encode_temporal_position(src_dense, lookup_index)
src_sp_embedding = self.encode_spatial_position(src_tp_embedding)
encoder_output = self.encoder(src_sp_embedding)
def encode(self, src_idx, src):
src_dense = self.encoder_dense(src[..., :1])
# tp_embedding = self.encode_temporal_position(src, src_idx)
# sp_embedding = self.encode_spatial_position(src)
week_embedding = self.temporal_section_week(src)
day_embedding = self.temporal_section_day(src)

embed = src_dense # + tp_embedding + sp_embedding
embed = paddle.concat([embed, week_embedding, day_embedding], axis=-1)

encoder_output = self.encoder(embed)
return encoder_output

def decode(self, tgt, encoder_output):
tgt_dense = self.decoder_dense(tgt)
tgt_tp_embedding = self.decode_temporal_position(tgt_dense)
tgt_sp_embedding = self.decode_spatial_position(tgt_tp_embedding)
decoder_output = self.decoder(x=tgt_sp_embedding, memory=encoder_output)
def decode(self, encoder_output, tgt_idx=None, tgt=None):
tgt_dense = self.decoder_dense(tgt[..., :1])
# tp_embedding = self.decode_temporal_position(tgt, tgt_idx)
# sp_embedding = self.decode_spatial_position(tgt)
week_embedding = self.temporal_section_week(tgt)
day_embedding = self.temporal_section_day(tgt)

embed = tgt_dense # + tp_embedding + sp_embedding
embed = paddle.concat([embed, week_embedding, day_embedding], axis=-1)

decoder_output = self.decoder(x=embed, memory=encoder_output)
return self.generator(decoder_output)

def forward(self, src, src_idx, tgt):
encoder_output = self.encode(src, src_idx)
output = self.decode(tgt, encoder_output)
def forward(self, src_idx, src, tgt_idx=None, tgt=None):
encoder_output = self.encode(src_idx, src)
output = self.decode(encoder_output, tgt_idx, tgt)
return output


class DecoderIndex(autograd.PyLayer):
@staticmethod
def forward(ctx, lags, his, his_span, interp_method="cubic"):
"""
计算给定输入序列的未来值,并返回计算结果。
传入lags, history,
计算序列位置对应位置的梯度, 并保存至backward

Args:
ctx (): 动态图计算上下文对象。
xde (): 未来值的输入序列, BaseXDE类型。
lags (paddle.Tensor): 用多少个过去的值来计算未来的这个值(未来值的滞后量)。
history (paddle.Tensor): 用于计算未来值的过去输入序列。
interp_method (str, optional): 插值方法,取值为 "linear"(线性插值),"cubic"(三次样条插值)或 "bez"(贝塞尔插值)。默认为 "linear"。

Returns:
paddle.Tensor: 计算结果,形状为 [batch_size, len_t, dims]。

Raises:
NotImplementedError: 如果interp_method不是上述三种情况之一, 将抛出NotImplementedError异常。
"""
with paddle.no_grad():
if interp_method == "linear":
interp = LinearInterpolation(his, his_span)
elif interp_method == "cubic":
interp = CubicHermiteSpline(his, his_span)
elif interp_method == "bez":
interp = BezierSpline(his, his_span)
else:
raise NotImplementedError

y_lags = interp.evaluate(lags)

derivative_lags = interp.derivative(lags)
ctx.save_for_backward(derivative_lags)

return y_lags

@staticmethod
def backward(ctx, grad_y):
# 计算history相应的梯度,并提取forward中保存的梯度,用于计算lag的梯度
# 在计算的过程中,无需更新history,仅更新lags即可
(derivative_lags,) = ctx.saved_tensor()
grad = grad_y * derivative_lags
grad = paddle.sum(grad, axis=[0, 1, 3])
return grad, None, None
# return None, grad_y_lags * derivative_lags, None, None, None


if __name__ == "__main__":
import os

# 将日志级别设置为6
os.environ["GLOG_v"] = "6"
import numpy as np
import paddle
import paddle.nn as nn
from args import args
from dataset import TrafficFlowDataset
from paddle.io import DataLoader
from paddle.nn.initializer import XavierUniform
from utils import get_adjacency_matrix_2direction, norm_adj_matrix

from paddlexde.functional import ddeint
from paddlexde.solver import Euler

default_dtype = paddle.get_default_dtype()
adj_matrix, _ = get_adjacency_matrix_2direction(args.adj_path, 80)
adj_matrix = paddle.to_tensor(norm_adj_matrix(adj_matrix), default_dtype)
Expand All @@ -123,9 +203,48 @@ def forward(self, src, src_idx, tgt):
nn.initializer.set_global_initializer(XavierUniform(), XavierUniform())
model = CorrSTN(args, adj_matrix, sc_matrix)

def collate_func(batch_data):
src_list, tgt_list = [], []

for item in batch_data:
if item[2]:
src_list.append(item[0])
tgt_list.append(item[1])

if len(src_list) == 0:
src_list.append(item[0])
tgt_list.append(item[1])

return paddle.stack(src_list), paddle.stack(tgt_list)

train_dataset = TrafficFlowDataset(args, "train")
train_dataloader = DataLoader(train_dataset, batch_size=args.batch_size)
his, tgt = next(iter(train_dataloader))
his_index = paddle.randint(shape=[36], high=his.shape[-2], low=0)
his_input = paddle.index_select(his, his_index, axis=-2)
output = model(his_input, his_index, tgt)
train_dataloader = DataLoader(
train_dataset, batch_size=args.batch_size, collate_fn=collate_func
)
src, tgt = next(iter(train_dataloader))
src_index = paddle.randint(shape=[12], high=src.shape[-2], low=0)
src_input = paddle.index_select(src, src_index, axis=-2)
decoder_input = paddle.concat([src[:, :, -1:, :], tgt[:, :, :-1, :]], axis=-2)

# 1. call model foreward (choose 1 or 2)
preds = model(src_index, src_input, None, decoder_input)
preds.backward()

# 2. call ddeint (choose 1 or 2)
y0 = decoder_input
preds = ddeint(
func=model,
y0=y0,
t_span=paddle.arange(args.tgt_len + 1),
lags=src_index,
his=src,
his_span=paddle.arange(args.his_len),
solver=Euler,
)
preds = preds[:, :, -args.tgt_len :, :]
preds.backward()

from paddleviz.paddleviz.viz import make_graph

dot = make_graph(preds, dpi="600")
dot.render("viz-result.gv", format="png", view=False)
Loading