diff --git a/README.md b/README.md new file mode 100644 index 0000000..4d4a51d --- /dev/null +++ b/README.md @@ -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】 diff --git a/timesead_ext/models/other/neutral_ad.py b/timesead_ext/models/other/neutral_ad.py index 3f1208b..b7c3bde 100644 --- a/timesead_ext/models/other/neutral_ad.py +++ b/timesead_ext/models/other/neutral_ad.py @@ -14,7 +14,7 @@ # You should have received a copy of the GNU Affero General Public License # along with this program. If not, see . -from typing import Tuple +from typing import Dict, List, Optional, Sequence, Tuple import math import torch @@ -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): @@ -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() @@ -169,7 +179,6 @@ 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'] @@ -177,47 +186,71 @@ def make_seq_nets(x_dim: int, config: dict): 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}') diff --git a/timesead_ext/models/transforms/__init__.py b/timesead_ext/models/transforms/__init__.py new file mode 100644 index 0000000..e729a9c --- /dev/null +++ b/timesead_ext/models/transforms/__init__.py @@ -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", +] diff --git a/timesead_ext/models/transforms/base.py b/timesead_ext/models/transforms/base.py new file mode 100644 index 0000000..9586b4a --- /dev/null +++ b/timesead_ext/models/transforms/base.py @@ -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) diff --git a/timesead_ext/models/transforms/freq.py b/timesead_ext/models/transforms/freq.py new file mode 100644 index 0000000..0826153 --- /dev/null +++ b/timesead_ext/models/transforms/freq.py @@ -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)] diff --git a/timesead_ext/models/transforms/group.py b/timesead_ext/models/transforms/group.py new file mode 100644 index 0000000..6f63aec --- /dev/null +++ b/timesead_ext/models/transforms/group.py @@ -0,0 +1,146 @@ +from __future__ import annotations + +from typing import Dict, Iterable, List, Tuple + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .base import Transform + + +def group_cfg( + k_group: int = 1, + enable_shift: bool = True, + enable_scale_bias: bool = True, + enable_warp: bool = True, + warp_knots: int = 8, + max_shift_frac: float = 0.1, + scale_range: float | Tuple[float, float] = 0.2, + bias_range: float | Tuple[float, float] = 0.1, + warp_strength: float = 1.0, +) -> Dict[str, object]: + return { + "k_group": k_group, + "enable_shift": enable_shift, + "enable_scale_bias": enable_scale_bias, + "enable_warp": enable_warp, + "warp_knots": warp_knots, + "max_shift_frac": max_shift_frac, + "scale_range": scale_range, + "bias_range": bias_range, + "warp_strength": warp_strength, + } + + +def _inverse_softplus(value: float) -> float: + return float(torch.log(torch.expm1(torch.tensor(value))).item()) + + +def _range_bounds(range_value: float | Iterable[float], center: float) -> Tuple[float, float]: + if isinstance(range_value, (tuple, list)): + low, high = range_value + return float(low), float(high) + span = float(range_value) + if center == 1.0: + if span <= 1.0: + return max(1.0 - span, 1e-4), 1.0 + span + return 1.0 / span, span + return center - span, center + span + + +class GroupTransform(Transform): + def __init__(self, channels: int, cfg: Dict[str, object]): + super().__init__() + self.channels = channels + self.enable_shift = bool(cfg.get("enable_shift", True)) + self.enable_scale_bias = bool(cfg.get("enable_scale_bias", True)) + self.enable_warp = bool(cfg.get("enable_warp", True)) + self.warp_knots = int(cfg.get("warp_knots", 8)) + self.max_shift_frac = float(cfg.get("max_shift_frac", 0.1)) + self.scale_range = cfg.get("scale_range", 0.2) + self.bias_range = cfg.get("bias_range", 0.1) + self.warp_strength = float(cfg.get("warp_strength", 1.0)) + + if self.enable_shift: + self.raw_shift = nn.Parameter(torch.zeros(1)) + else: + self.register_parameter("raw_shift", None) + + if self.enable_scale_bias: + self.raw_scale = nn.Parameter(torch.zeros(channels)) + self.raw_bias = nn.Parameter(torch.zeros(channels)) + else: + self.register_parameter("raw_scale", None) + self.register_parameter("raw_bias", None) + + if self.enable_warp: + if self.warp_knots < 2: + raise ValueError("warp_knots must be >= 2") + knot_init = _inverse_softplus(1.0) + self.raw_warp = nn.Parameter(torch.full((self.warp_knots,), knot_init)) + else: + self.register_parameter("raw_warp", None) + + self._scale_offset = _inverse_softplus(1.0) + + def _base_coords(self, length: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor: + return torch.linspace(-1.0, 1.0, length, device=device, dtype=dtype) + + def _warp_coords(self, length: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor: + base = self._base_coords(length, device, dtype) + if not self.enable_warp: + return base + deltas = F.softplus(self.raw_warp) + cumulative = torch.cumsum(deltas, dim=0) + cumulative = (cumulative - cumulative[0]) / (cumulative[-1] - cumulative[0]) + warp = cumulative * 2.0 - 1.0 + warp = warp.view(1, 1, -1) + warp = F.interpolate(warp, size=length, mode="linear", align_corners=True) + warp = warp.view(-1) + strength = max(0.0, min(self.warp_strength, 1.0)) + return base + strength * (warp - base) + + def _shift_coords(self, coords: torch.Tensor) -> torch.Tensor: + if not self.enable_shift: + return coords + shift = self.max_shift_frac * torch.tanh(self.raw_shift)[0] + return coords + shift * 2.0 + + def _apply_time_transform(self, x: torch.Tensor) -> torch.Tensor: + if not (self.enable_shift or self.enable_warp): + return x + length = x.shape[-1] + coords = self._warp_coords(length, x.device, x.dtype) + coords = self._shift_coords(coords) + grid = torch.zeros((x.shape[0], 1, length, 2), device=x.device, dtype=x.dtype) + grid[..., 0] = coords.view(1, 1, -1).expand(x.shape[0], 1, -1) + return F.grid_sample( + x.unsqueeze(2), + grid, + mode="bilinear", + padding_mode="border", + align_corners=True, + ).squeeze(2) + + def _apply_scale_bias(self, x: torch.Tensor) -> torch.Tensor: + if not self.enable_scale_bias: + return x + scale = F.softplus(self.raw_scale + self._scale_offset) + min_scale, max_scale = _range_bounds(self.scale_range, 1.0) + if min_scale != max_scale: + scale = min_scale + (max_scale - min_scale) * torch.sigmoid(scale - 1.0) + bias_min, bias_max = _range_bounds(self.bias_range, 0.0) + bias = bias_min + (bias_max - bias_min) * torch.sigmoid(self.raw_bias) + return x * scale.view(1, -1, 1) + bias.view(1, -1, 1) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self._apply_time_transform(x) + return self._apply_scale_bias(x) + + +def make_group_family(channels: int, cfg: Dict[str, object]) -> List[GroupTransform]: + k_group = int(cfg.get("k_group", 1)) + if k_group < 1: + raise ValueError("k_group must be >= 1") + return [GroupTransform(channels, cfg) for _ in range(k_group)] diff --git a/timesead_ext/models/transforms/invertible.py b/timesead_ext/models/transforms/invertible.py new file mode 100644 index 0000000..f52f1d7 --- /dev/null +++ b/timesead_ext/models/transforms/invertible.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +from typing import Dict, List, Tuple + +import torch +import torch.nn as nn + +from .base import Transform + + +def invertible_cfg( + num_flows: int = 4, + hidden: int = 32, + kernel_size: int = 3, + clamp: float = 2.0, + share_across_k: bool = False, + k_invertible: int = 1, +) -> Dict[str, object]: + return { + "num_flows": num_flows, + "hidden": hidden, + "kernel_size": kernel_size, + "clamp": clamp, + "share_across_k": share_across_k, + "k_invertible": k_invertible, + } + + +class AffineCoupling(nn.Module): + def __init__(self, channels: int, hidden: int, kernel_size: int, clamp: float, swap: bool = False): + super().__init__() + self.clamp = clamp + self.swap = swap + c1 = channels // 2 + c2 = channels - c1 + self.c1 = c1 + self.c2 = c2 + if swap: + cond_channels = c2 + target_channels = c1 + else: + cond_channels = c1 + target_channels = c2 + padding = kernel_size // 2 + self.net = nn.Sequential( + nn.Conv1d(cond_channels, hidden, kernel_size, padding=padding), + nn.ReLU(inplace=True), + nn.Conv1d(hidden, target_channels * 2, kernel_size, padding=padding), + ) + + def _split(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + x1, x2 = x.split([self.c1, self.c2], dim=1) + return x1, x2 + + def _condition(self, cond: torch.Tensor) -> torch.Tensor: + scale_shift = self.net(cond) + scale, shift = scale_shift.chunk(2, dim=1) + scale = self.clamp * torch.tanh(scale) + return scale, shift + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x1, x2 = self._split(x) + if self.swap: + scale, shift = self._condition(x2) + y1 = x1 * torch.exp(scale) + shift + return torch.cat([y1, x2], dim=1) + scale, shift = self._condition(x1) + y2 = x2 * torch.exp(scale) + shift + return torch.cat([x1, y2], dim=1) + + def inverse(self, y: torch.Tensor) -> torch.Tensor: + y1, y2 = self._split(y) + if self.swap: + scale, shift = self._condition(y2) + x1 = (y1 - shift) * torch.exp(-scale) + return torch.cat([x1, y2], dim=1) + scale, shift = self._condition(y1) + x2 = (y2 - shift) * torch.exp(-scale) + return torch.cat([y1, x2], dim=1) + + +class InvertibleFlow(Transform): + def __init__(self, channels: int, num_flows: int, hidden: int, kernel_size: int, clamp: float): + super().__init__() + self.layers = nn.ModuleList( + [ + AffineCoupling( + channels=channels, + hidden=hidden, + kernel_size=kernel_size, + clamp=clamp, + swap=bool(idx % 2), + ) + for idx in range(num_flows) + ] + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + for layer in self.layers: + x = layer(x) + return x + + def inverse(self, y: torch.Tensor) -> torch.Tensor: + for layer in reversed(self.layers): + y = layer.inverse(y) + return y + + +def make_invertible_family(channels: int, cfg: Dict[str, object]) -> List[InvertibleFlow]: + num_flows = int(cfg["num_flows"]) + hidden = int(cfg["hidden"]) + kernel_size = int(cfg["kernel_size"]) + clamp = float(cfg["clamp"]) + k_invertible = int(cfg.get("k_invertible", 1)) + share_across_k = bool(cfg.get("share_across_k", False)) + + if k_invertible < 1: + raise ValueError("k_invertible must be >= 1") + + if share_across_k: + flow = InvertibleFlow(channels, num_flows, hidden, kernel_size, clamp) + return [flow for _ in range(k_invertible)] + + return [InvertibleFlow(channels, num_flows, hidden, kernel_size, clamp) for _ in range(k_invertible)]