Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
23 changes: 23 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
# TimeSeAD Extensions

This repository adds optional transform families for the NeutralAD model in `timesead_ext`.

## Optional transform families

NeutralAD can now include three optional transform families in its transform bank:

- **Invertible flows**: stacks of affine coupling layers that learn invertible, channel-wise transformations for time series data. Enable with `use_invertible_transforms=True`. Configure with `invertible_cfg(...)` (e.g., number of flows, hidden size, kernel size).【F:timesead_ext/models/transforms/invertible.py†L1-L118】【F:timesead_ext/models/other/neutral_ad.py†L196-L223】
- **Group transforms**: learnable time-domain warping, shifting, and per-channel scale/bias adjustments to create structured augmentations. Enable with `use_group_transforms=True`. Configure with `group_cfg(...)` (e.g., shift/warp toggles and ranges).【F:timesead_ext/models/transforms/group.py†L1-L146】【F:timesead_ext/models/other/neutral_ad.py†L196-L223】
- **Frequency-orthogonal transforms**: apply orthogonal mixing in the frequency domain (either across channels or frequency blocks) via Cayley-parameterized skew matrices. Enable with `use_freq_ortho_transforms=True`. Configure with `freq_cfg(...)` (e.g., mode, block size).【F:timesead_ext/models/transforms/freq.py†L1-L127】【F:timesead_ext/models/other/neutral_ad.py†L196-L223】

## New configuration flags

NeutralAD accepts new flags for assembling the transform bank:

- `use_invertible_transforms`: include the invertible flow family.
- `use_group_transforms`: include the group transform family.
- `use_freq_ortho_transforms`: include the frequency-orthogonal family.
- `keep_base_transforms`: keep the original `SeqTransformNet` transforms (default `True`). Set to `False` if you want only the optional families.
- `transform_families`: provide additional custom transform families (list of lists) to append after the built-ins.

These flags can be combined, and at least one transform must be configured.【F:timesead_ext/models/other/neutral_ad.py†L196-L229】
69 changes: 51 additions & 18 deletions timesead_ext/models/other/neutral_ad.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@
# You should have received a copy of the GNU Affero General Public License
# along with this program. If not, see <https://www.gnu.org/licenses/>.

from typing import Tuple
from typing import Dict, List, Optional, Sequence, Tuple

import math
import torch
Expand All @@ -25,6 +25,16 @@
from timesead.models import BaseModel
from timesead.optim.loss import Loss
from timesead.utils.utils import pack_tuple
from timesead_ext.models.transforms import (
Transform,
TransformBank,
freq_cfg as default_freq_cfg,
group_cfg as default_group_cfg,
invertible_cfg as default_invertible_cfg,
make_freq_family,
make_group_family,
make_invertible_family,
)


class ResTrans1DBlock(torch.nn.Module):
Expand Down Expand Up @@ -59,7 +69,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
return out


class SeqTransformNet(nn.Module):
class SeqTransformNet(Transform):
def __init__(self, x_dim: int, hdim: int, num_layers: int):
super().__init__()
self.relu = nn.ReLU()
Expand Down Expand Up @@ -169,55 +179,78 @@ def make_seq_nets(x_dim: int, config: dict):
enc_hdim = config['enc_hdim']
z_dim = config['latent_dim']
x_len = config['x_length']
trans_nlayers = config['trans_nlayers']
num_trans = config['num_trans']
batch_norm = config['batch_norm']

enc = nn.ModuleList([
SeqEncoder(x_dim, x_len, enc_hdim, z_dim, config['enc_bias'], enc_nlayers, batch_norm)
for _ in range(num_trans + 1)
])
trans = nn.ModuleList([
SeqTransformNet(x_dim, x_dim, trans_nlayers) for _ in range(num_trans)
])

return enc, trans
return enc


class NeutralAD(BaseModel):
def __init__(self, ts_channels: int, seq_len: int, num_trans: int = 4, trans_type: str = 'residual',
enc_hdim: int = 32, enc_nlayers: int = 4, trans_nlayers: int = 4, latent_dim: int = 32,
batch_norm: bool = False, enc_bias: bool = False):
batch_norm: bool = False, enc_bias: bool = False,
transform_families: Optional[Sequence[Sequence[Transform]]] = None,
use_invertible_transforms: bool = False,
use_group_transforms: bool = False,
use_freq_ortho_transforms: bool = False,
keep_base_transforms: bool = True,
invertible_cfg: Optional[Dict[str, object]] = None,
group_cfg: Optional[Dict[str, object]] = None,
freq_cfg: Optional[Dict[str, object]] = None):
super().__init__()

assert num_trans > 1, 'num_trans must be > 1'
self.num_trans = num_trans
self.trans_type = trans_type
self.z_dim = latent_dim
transforms: List[Transform] = []
if keep_base_transforms:
if num_trans < 1:
raise ValueError('num_trans must be >= 1 when keep_base_transforms is True')
transforms.extend(
[SeqTransformNet(ts_channels, ts_channels, trans_nlayers) for _ in range(num_trans)]
)
if use_invertible_transforms:
invertible_cfg_data = invertible_cfg or default_invertible_cfg()
transforms.extend(make_invertible_family(ts_channels, invertible_cfg_data))
if use_group_transforms:
group_cfg_data = group_cfg or default_group_cfg()
transforms.extend(make_group_family(ts_channels, group_cfg_data))
if use_freq_ortho_transforms:
freq_cfg_data = freq_cfg or default_freq_cfg()
transforms.extend(make_freq_family(ts_channels, freq_cfg_data))
if transform_families:
for family in transform_families:
transforms.extend(family)
if not transforms:
raise ValueError('At least one transform must be configured for NeutralAD')
self.transform_bank = TransformBank(transforms)
self.num_trans = len(self.transform_bank)
config = dict(
enc_nlayers=enc_nlayers,
enc_hdim=enc_hdim,
latent_dim=latent_dim,
x_length=seq_len,
trans_nlayers=trans_nlayers,
num_trans=num_trans,
num_trans=self.num_trans,
batch_norm=batch_norm,
enc_bias=enc_bias
)
self.enc, self.trans = make_seq_nets(ts_channels, config)
self.enc = make_seq_nets(ts_channels, config)

def forward(self, inputs: Tuple[torch.Tensor, ...]) -> torch.Tensor:
x, = inputs
x = x.float()
x = x.permute(0, 2, 1)

masks = [self.trans[i](x) for i in range(self.num_trans)]
masks = self.transform_bank(x)
if self.trans_type == 'forward':
x_t = torch.stack(masks, dim=1)
x_t = masks
elif self.trans_type == 'mul':
x_t = torch.stack([torch.sigmoid(mask) for mask in masks], dim=1) * x.unsqueeze(1)
x_t = torch.sigmoid(masks) * x.unsqueeze(1)
elif self.trans_type == 'residual':
x_t = torch.stack(masks, dim=1) + x.unsqueeze(1)
x_t = masks + x.unsqueeze(1)
else:
raise ValueError(f'Unknown trans_type: {self.trans_type}')

Expand Down
18 changes: 18 additions & 0 deletions timesead_ext/models/transforms/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
from .base import Transform, TransformBank
from .freq import FreqTransform, freq_cfg, make_freq_family
from .group import GroupTransform, group_cfg, make_group_family
from .invertible import InvertibleFlow, invertible_cfg, make_invertible_family

__all__ = [
"Transform",
"TransformBank",
"FreqTransform",
"freq_cfg",
"make_freq_family",
"GroupTransform",
"group_cfg",
"make_group_family",
"InvertibleFlow",
"invertible_cfg",
"make_invertible_family",
]
22 changes: 22 additions & 0 deletions timesead_ext/models/transforms/base.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
from typing import Iterable, List

import torch
import torch.nn as nn


class Transform(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor: # pragma: no cover - interface
raise NotImplementedError


class TransformBank(nn.Module):
def __init__(self, transforms: Iterable[Transform]):
super().__init__()
self.transforms = nn.ModuleList(list(transforms))

def __len__(self) -> int:
return len(self.transforms)

def forward(self, x: torch.Tensor) -> torch.Tensor:
outputs: List[torch.Tensor] = [transform(x) for transform in self.transforms]
return torch.stack(outputs, dim=1)
127 changes: 127 additions & 0 deletions timesead_ext/models/transforms/freq.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,127 @@
from __future__ import annotations

import math
from typing import Dict, List

import torch
import torch.nn as nn

from .base import Transform


def freq_cfg(
k_freq: int = 1,
mode: str = "channel",
freq_block: int = 8,
init_identity: bool = True,
eps: float = 1e-5,
) -> Dict[str, object]:
return {
"k_freq": k_freq,
"mode": mode,
"freq_block": freq_block,
"init_identity": init_identity,
"eps": eps,
}


def _cayley(skew: torch.Tensor, eps: float) -> torch.Tensor:
eye = torch.eye(skew.shape[-1], device=skew.device, dtype=skew.dtype)
eye = eye.expand(skew.shape[:-2] + eye.shape)
a = eye + skew
b = eye - skew
if eps > 0:
b = b + eps * eye
return torch.linalg.solve(b, a)


class FreqTransform(Transform):
def __init__(self, channels: int, cfg: Dict[str, object]):
super().__init__()
self.channels = channels
self.mode = str(cfg.get("mode", "channel")).lower()
self.freq_block = int(cfg.get("freq_block", 8))
self.init_identity = bool(cfg.get("init_identity", True))
self.eps = float(cfg.get("eps", 1e-5))

if self.freq_block < 1:
raise ValueError("freq_block must be >= 1")
if self.mode not in {"channel", "freq"}:
raise ValueError("mode must be 'channel' or 'freq'")

self.raw_skew: nn.Parameter | None = None
self._freq_bins: int | None = None

def _init_params(self, freq_bins: int, device: torch.device, dtype: torch.dtype) -> None:
if self.mode == "channel":
shape = (freq_bins, self.channels, self.channels)
else:
block_count = math.ceil(freq_bins / self.freq_block)
shape = (self.channels, block_count, self.freq_block, self.freq_block)

if self.init_identity:
param = torch.zeros(shape, device=device, dtype=dtype)
else:
param = 0.01 * torch.randn(shape, device=device, dtype=dtype)

self.raw_skew = nn.Parameter(param)
self._freq_bins = freq_bins

def _ensure_params(self, freq_bins: int, device: torch.device, dtype: torch.dtype) -> None:
if self.raw_skew is None:
self._init_params(freq_bins, device, dtype)
return
if self._freq_bins != freq_bins:
raise ValueError(
f"Input frequency bins changed from {self._freq_bins} to {freq_bins}. "
"Recreate the transform for a new input length."
)

def _mix_channels(self, x_freq: torch.Tensor) -> torch.Tensor:
assert self.raw_skew is not None
skew = self.raw_skew - self.raw_skew.transpose(-1, -2)
ortho = _cayley(skew, self.eps)
freq_t = x_freq.permute(0, 2, 1)
mixed = torch.einsum("fij,bfj->bfi", ortho, freq_t)
return mixed.permute(0, 2, 1)

def _mix_freq_blocks(self, x_freq: torch.Tensor) -> torch.Tensor:
assert self.raw_skew is not None
skew = self.raw_skew - self.raw_skew.transpose(-1, -2)
ortho = _cayley(skew, self.eps)

batch, channels, freq_bins = x_freq.shape
block_count = ortho.shape[1]
output = x_freq.clone()

for block_idx in range(block_count):
start = block_idx * self.freq_block
if start >= freq_bins:
break
end = min(start + self.freq_block, freq_bins)
block = x_freq[:, :, start:end]
if end - start < self.freq_block:
pad = self.freq_block - (end - start)
block = torch.nn.functional.pad(block, (0, pad))
mixed = torch.einsum("cij,bcj->bci", ortho[:, block_idx], block)
output[:, :, start:end] = mixed[:, :, : end - start]

return output

def forward(self, x: torch.Tensor) -> torch.Tensor:
freq = torch.fft.rfft(x, dim=-1)
self._ensure_params(freq.shape[-1], freq.device, freq.dtype)

if self.mode == "channel":
mixed = self._mix_channels(freq)
else:
mixed = self._mix_freq_blocks(freq)

return torch.fft.irfft(mixed, n=x.shape[-1], dim=-1)


def make_freq_family(channels: int, cfg: Dict[str, object]) -> List[FreqTransform]:
k_freq = int(cfg.get("k_freq", 1))
if k_freq < 1:
raise ValueError("k_freq must be >= 1")
return [FreqTransform(channels, cfg) for _ in range(k_freq)]
Loading