minor changes
This commit is contained in:
@@ -43,10 +43,7 @@ class RMSNorm(nnx.Module):
|
|||||||
self.num_features = num_features
|
self.num_features = num_features
|
||||||
self.epsilon = epsilon
|
self.epsilon = epsilon
|
||||||
|
|
||||||
def __call__(
|
def __call__(self, inputs: Float[Array, "b ... d"]) -> Float[Array, "b ... d"]:
|
||||||
self, inputs: Float[Array, "b ... d"]
|
|
||||||
) -> Float[Array, "b ... d"]:
|
|
||||||
|
|
||||||
var = jnp.mean(jnp.square(inputs), axis=-1, keepdims=True)
|
var = jnp.mean(jnp.square(inputs), axis=-1, keepdims=True)
|
||||||
normed_inputs = inputs * jax.lax.rsqrt(var + self.epsilon)
|
normed_inputs = inputs * jax.lax.rsqrt(var + self.epsilon)
|
||||||
normed_inputs *= self.scale
|
normed_inputs *= self.scale
|
||||||
@@ -58,18 +55,14 @@ class LayerNorm(nnx.Module):
|
|||||||
|
|
||||||
__data__ = ("scale", "bias")
|
__data__ = ("scale", "bias")
|
||||||
|
|
||||||
def __init__(
|
def __init__(self, num_features: int, *, epsilon: float = 1e-6, rngs=nnx.Rngs(42)):
|
||||||
self, num_features: int, *, epsilon: float = 1e-6, rngs=nnx.Rngs(42)
|
|
||||||
):
|
|
||||||
del rngs
|
del rngs
|
||||||
self.scale = nnx.Param(jnp.ones(shape=(num_features,)))
|
self.scale = nnx.Param(jnp.ones(shape=(num_features,)))
|
||||||
self.bias = nnx.Param(jnp.zeros(shape=(num_features,)))
|
self.bias = nnx.Param(jnp.zeros(shape=(num_features,)))
|
||||||
self.num_features = num_features
|
self.num_features = num_features
|
||||||
self.epsilon = epsilon
|
self.epsilon = epsilon
|
||||||
|
|
||||||
def __call__(
|
def __call__(self, inputs: Float[Array, "b ... d"]) -> Float[Array, "b ... d"]:
|
||||||
self, inputs: Float[Array, "b ... d"]
|
|
||||||
) -> Float[Array, "b ... d"]:
|
|
||||||
mean = jnp.mean(inputs, axis=-1, keepdims=True)
|
mean = jnp.mean(inputs, axis=-1, keepdims=True)
|
||||||
var = jnp.mean(jnp.square(inputs - mean), axis=-1, keepdims=True)
|
var = jnp.mean(jnp.square(inputs - mean), axis=-1, keepdims=True)
|
||||||
normed_inputs = (inputs - mean) * jax.lax.rsqrt(var + self.epsilon)
|
normed_inputs = (inputs - mean) * jax.lax.rsqrt(var + self.epsilon)
|
||||||
|
|||||||
@@ -91,8 +91,7 @@ class RotaryPositionalEmbedding(nnx.Module):
|
|||||||
half_embedding_dim = self.embedding_dims // 2
|
half_embedding_dim = self.embedding_dims // 2
|
||||||
fraction = 2 * jnp.arange(0, half_embedding_dim) / self.embedding_dims
|
fraction = 2 * jnp.arange(0, half_embedding_dim) / self.embedding_dims
|
||||||
timescale = (
|
timescale = (
|
||||||
self.min_timescale
|
self.min_timescale * (self.max_timescale / self.min_timescale) ** fraction
|
||||||
* (self.max_timescale / self.min_timescale) ** fraction
|
|
||||||
)
|
)
|
||||||
if position is None:
|
if position is None:
|
||||||
seq_length = inputs.shape[1]
|
seq_length = inputs.shape[1]
|
||||||
@@ -128,9 +127,7 @@ class PerDimScale(nnx.Module):
|
|||||||
|
|
||||||
def __call__(self, x: Float[Array, "b ... d"]) -> Float[Array, "b ... d"]:
|
def __call__(self, x: Float[Array, "b ... d"]) -> Float[Array, "b ... d"]:
|
||||||
return x * (
|
return x * (
|
||||||
1.442695041
|
1.442695041 / jnp.sqrt(self.num_dims) * jax.nn.softplus(self.per_dim_scale)
|
||||||
/ jnp.sqrt(self.num_dims)
|
|
||||||
* jax.nn.softplus(self.per_dim_scale)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -232,9 +229,7 @@ class MultiHeadAttention(nnx.Module):
|
|||||||
value = self.value(inputs_q)
|
value = self.value(inputs_q)
|
||||||
|
|
||||||
if decode_cache is None:
|
if decode_cache is None:
|
||||||
num_masked = jnp.sum(
|
num_masked = jnp.sum(patch_mask.astype(jnp.int32), axis=-1, keepdims=False)
|
||||||
patch_mask.astype(jnp.int32), axis=-1, keepdims=False
|
|
||||||
)
|
|
||||||
next_index = jnp.zeros_like(num_masked, dtype=jnp.int32)
|
next_index = jnp.zeros_like(num_masked, dtype=jnp.int32)
|
||||||
else:
|
else:
|
||||||
num_masked = (
|
num_masked = (
|
||||||
@@ -277,9 +272,7 @@ class MultiHeadAttention(nnx.Module):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Training
|
# Training
|
||||||
attn_mask = make_attn_mask(
|
attn_mask = make_attn_mask(query_length=n_patches, num_all_masked_kv=num_masked)
|
||||||
query_length=n_patches, num_all_masked_kv=num_masked
|
|
||||||
)
|
|
||||||
|
|
||||||
# apply attention
|
# apply attention
|
||||||
x = self.attention_fn(
|
x = self.attention_fn(
|
||||||
@@ -357,9 +350,7 @@ class Transformer(nnx.Module):
|
|||||||
)
|
)
|
||||||
attn_output = self.post_attn_ln(attn_output) + input_embeddings
|
attn_output = self.post_attn_ln(attn_output) + input_embeddings
|
||||||
output_embeddings = (
|
output_embeddings = (
|
||||||
self.post_ff_ln(
|
self.post_ff_ln(self.ff1(self.activation(self.ff0(self.pre_ff_ln(attn_output)))))
|
||||||
self.ff1(self.activation(self.ff0(self.pre_ff_ln(attn_output))))
|
|
||||||
)
|
|
||||||
+ attn_output
|
+ attn_output
|
||||||
)
|
)
|
||||||
return output_embeddings, decode_cache
|
return output_embeddings, decode_cache
|
||||||
|
|||||||
@@ -19,7 +19,6 @@ import functools
|
|||||||
import jax
|
import jax
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
import jaxtyping
|
import jaxtyping
|
||||||
import typeguard
|
|
||||||
|
|
||||||
Float = jaxtyping.Float
|
Float = jaxtyping.Float
|
||||||
Array = jaxtyping.Array
|
Array = jaxtyping.Array
|
||||||
|
|||||||
@@ -19,7 +19,9 @@ import functools
|
|||||||
import gc
|
import gc
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
from typing import Any, Callable
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Callable, Dict
|
||||||
|
|
||||||
import einshape
|
import einshape
|
||||||
from flax import nnx
|
from flax import nnx
|
||||||
@@ -127,9 +129,7 @@ class TimesFM_2p5_200M_flax_module(nnx.Module): # pylint: disable=invalid-name
|
|||||||
masks: Bool[Array, "b n p"],
|
masks: Bool[Array, "b n p"],
|
||||||
decode_cache: util.DecodeCache | None = None,
|
decode_cache: util.DecodeCache | None = None,
|
||||||
):
|
):
|
||||||
tokenizer_inputs = jnp.concatenate(
|
tokenizer_inputs = jnp.concatenate([inputs, masks.astype(inputs.dtype)], axis=-1)
|
||||||
[inputs, masks.astype(inputs.dtype)], axis=-1
|
|
||||||
)
|
|
||||||
input_embeddings = self.tokenizer(tokenizer_inputs)
|
input_embeddings = self.tokenizer(tokenizer_inputs)
|
||||||
if decode_cache is None:
|
if decode_cache is None:
|
||||||
decode_cache = [None] * self.x
|
decode_cache = [None] * self.x
|
||||||
@@ -164,16 +164,10 @@ class TimesFM_2p5_200M_flax_module(nnx.Module): # pylint: disable=invalid-name
|
|||||||
decode_cache = util.DecodeCache(
|
decode_cache = util.DecodeCache(
|
||||||
next_index=jnp.zeros(shape=(self.x, batch_size), dtype=jnp.int32),
|
next_index=jnp.zeros(shape=(self.x, batch_size), dtype=jnp.int32),
|
||||||
num_masked=jnp.zeros(shape=(self.x, batch_size), dtype=jnp.int32),
|
num_masked=jnp.zeros(shape=(self.x, batch_size), dtype=jnp.int32),
|
||||||
key=jnp.zeros(
|
key=jnp.zeros(shape=(self.x, batch_size, decode_cache_size, self.h, self.hd)),
|
||||||
shape=(self.x, batch_size, decode_cache_size, self.h, self.hd)
|
value=jnp.zeros(shape=(self.x, batch_size, decode_cache_size, self.h, self.hd)),
|
||||||
),
|
|
||||||
value=jnp.zeros(
|
|
||||||
shape=(self.x, batch_size, decode_cache_size, self.h, self.hd)
|
|
||||||
),
|
|
||||||
)
|
|
||||||
normed_inputs = revin(
|
|
||||||
patched_inputs, context_mu, context_sigma, reverse=False
|
|
||||||
)
|
)
|
||||||
|
normed_inputs = revin(patched_inputs, context_mu, context_sigma, reverse=False)
|
||||||
normed_inputs = jnp.where(patched_masks, 0.0, normed_inputs)
|
normed_inputs = jnp.where(patched_masks, 0.0, normed_inputs)
|
||||||
(_, _, normed_outputs, normed_quantile_spread), decode_cache = self(
|
(_, _, normed_outputs, normed_quantile_spread), decode_cache = self(
|
||||||
normed_inputs, patched_masks, decode_cache
|
normed_inputs, patched_masks, decode_cache
|
||||||
@@ -205,9 +199,7 @@ class TimesFM_2p5_200M_flax_module(nnx.Module): # pylint: disable=invalid-name
|
|||||||
xs=(new_patched_input, new_mask),
|
xs=(new_patched_input, new_mask),
|
||||||
axis=1,
|
axis=1,
|
||||||
)
|
)
|
||||||
new_normed_input = revin(
|
new_normed_input = revin(new_patched_input, new_mu, new_sigma, reverse=False)
|
||||||
new_patched_input, new_mu, new_sigma, reverse=False
|
|
||||||
)
|
|
||||||
(_, _, new_normed_output, _), decode_cache = module(
|
(_, _, new_normed_output, _), decode_cache = module(
|
||||||
new_normed_input, new_mask, decode_cache
|
new_normed_input, new_mask, decode_cache
|
||||||
)
|
)
|
||||||
@@ -298,15 +290,11 @@ def _force_flip_invariance_fn(
|
|||||||
flipped_pf_outputs = _flip_quantile_fn(flipped_pf_outputs)
|
flipped_pf_outputs = _flip_quantile_fn(flipped_pf_outputs)
|
||||||
flipped_pf_outputs = jax_einshape("tb...->(tb)...", flipped_pf_outputs)
|
flipped_pf_outputs = jax_einshape("tb...->(tb)...", flipped_pf_outputs)
|
||||||
flipped_quantile_spreads = _flip_quantile_fn(flipped_quantile_spreads)
|
flipped_quantile_spreads = _flip_quantile_fn(flipped_quantile_spreads)
|
||||||
flipped_quantile_spreads = jax_einshape(
|
flipped_quantile_spreads = jax_einshape("tb...->(tb)...", flipped_quantile_spreads)
|
||||||
"tb...->(tb)...", flipped_quantile_spreads
|
|
||||||
)
|
|
||||||
to_concat = [flipped_pf_outputs[:, -1, ...]]
|
to_concat = [flipped_pf_outputs[:, -1, ...]]
|
||||||
if flipped_ar_outputs is not None:
|
if flipped_ar_outputs is not None:
|
||||||
flipped_ar_outputs = _flip_quantile_fn(flipped_ar_outputs)
|
flipped_ar_outputs = _flip_quantile_fn(flipped_ar_outputs)
|
||||||
flipped_ar_outputs = jax_einshape(
|
flipped_ar_outputs = jax_einshape("tbno...->(tb)(no)...", flipped_ar_outputs)
|
||||||
"tbno...->(tb)(no)...", flipped_ar_outputs
|
|
||||||
)
|
|
||||||
to_concat.append(flipped_ar_outputs)
|
to_concat.append(flipped_ar_outputs)
|
||||||
flipped_full_forecast = jnp.concatenate(to_concat, axis=1)
|
flipped_full_forecast = jnp.concatenate(to_concat, axis=1)
|
||||||
|
|
||||||
@@ -318,9 +306,7 @@ def _force_flip_invariance_fn(
|
|||||||
static_argnames=("max_horizon",),
|
static_argnames=("max_horizon",),
|
||||||
donate_argnums=(0,),
|
donate_argnums=(0,),
|
||||||
)
|
)
|
||||||
def _use_continuous_quantile_head_fn(
|
def _use_continuous_quantile_head_fn(full_forecast, quantile_spreads, max_horizon):
|
||||||
full_forecast, quantile_spreads, max_horizon
|
|
||||||
):
|
|
||||||
"""Uses continuous quantile head."""
|
"""Uses continuous quantile head."""
|
||||||
to_stack = [full_forecast[..., :max_horizon, 0]]
|
to_stack = [full_forecast[..., :max_horizon, 0]]
|
||||||
for quantile_index in [1, 2, 3, 4]:
|
for quantile_index in [1, 2, 3, 4]:
|
||||||
@@ -464,31 +450,48 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(
|
def from_pretrained(
|
||||||
cls,
|
cls,
|
||||||
|
model_id: str = "google/timesfm-2.5-200m-flax",
|
||||||
*,
|
*,
|
||||||
path: str | None = None,
|
revision: str | None = None,
|
||||||
hf_repo_id: str | None = "google/timesfm-2.5-200m-flax",
|
cache_dir: str | Path | None = None,
|
||||||
|
force_download: bool = False,
|
||||||
|
proxies: Dict | None = None,
|
||||||
|
resume_download: bool | None = None,
|
||||||
|
local_files_only: bool | None = None,
|
||||||
|
token: str | None = None,
|
||||||
|
**model_kwargs,
|
||||||
):
|
):
|
||||||
"""Loads a Flax TimesFM model."""
|
"""Loads a Flax TimesFM model."""
|
||||||
if path:
|
|
||||||
pass
|
|
||||||
elif hf_repo_id:
|
|
||||||
logging.info(
|
|
||||||
"Downloading checkpoint from HuggingFace repo %s", hf_repo_id
|
|
||||||
)
|
|
||||||
path = huggingface_hub.snapshot_download(hf_repo_id)
|
|
||||||
logging.info("Loading checkpoint from: %s", path)
|
|
||||||
else:
|
|
||||||
raise ValueError("Either path or hf_repo_id must be provided.")
|
|
||||||
|
|
||||||
instance = cls()
|
# Create an instance of the model wrapper class.
|
||||||
|
instance = cls(**model_kwargs)
|
||||||
|
|
||||||
|
# Determine the path to the model weights.
|
||||||
|
model_file_path = ""
|
||||||
|
if os.path.isdir(model_id):
|
||||||
|
logging.info("Loading checkpoint from local directory: %s", model_id)
|
||||||
|
model_file_path = model_id
|
||||||
|
else:
|
||||||
|
logging.info("Downloading checkpoint from Hugging Face repo %s", model_id)
|
||||||
|
model_file_path = huggingface_hub.snapshot_download(
|
||||||
|
repo_id=model_id,
|
||||||
|
revision=revision,
|
||||||
|
cache_dir=cache_dir,
|
||||||
|
force_download=force_download,
|
||||||
|
proxies=proxies,
|
||||||
|
resume_download=resume_download,
|
||||||
|
token=token,
|
||||||
|
local_files_only=local_files_only,
|
||||||
|
)
|
||||||
|
logging.info("Loading checkpoint from: %s", model_file_path)
|
||||||
|
|
||||||
checkpointer = ocp.StandardCheckpointer()
|
checkpointer = ocp.StandardCheckpointer()
|
||||||
graph, state = nnx.split(instance.model)
|
graph, state = nnx.split(instance.model)
|
||||||
state = checkpointer.restore(path, state)
|
state = checkpointer.restore(model_file_path, state)
|
||||||
instance.model = nnx.merge(graph, state)
|
instance.model = nnx.merge(graph, state)
|
||||||
return instance
|
return instance
|
||||||
|
|
||||||
def compile(self, forecast_config: configs.ForecastConfig, **kwargs):
|
def compile(self, forecast_config: configs.ForecastConfig, **kwargs):
|
||||||
|
|
||||||
# Acrobym used during validation.
|
# Acrobym used during validation.
|
||||||
fc = forecast_config
|
fc = forecast_config
|
||||||
if fc.max_context % self.model.p != 0:
|
if fc.max_context % self.model.p != 0:
|
||||||
@@ -496,8 +499,7 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
|
|||||||
"When compiling, max context needs to be multiple of the patch size"
|
"When compiling, max context needs to be multiple of the patch size"
|
||||||
" %d. Using max context = %d instead.",
|
" %d. Using max context = %d instead.",
|
||||||
self.model.p,
|
self.model.p,
|
||||||
new_context := math.ceil(fc.max_context / self.model.p)
|
new_context := math.ceil(fc.max_context / self.model.p) * self.model.p,
|
||||||
* self.model.p,
|
|
||||||
)
|
)
|
||||||
fc = dataclasses.replace(fc, max_context=new_context)
|
fc = dataclasses.replace(fc, max_context=new_context)
|
||||||
if fc.max_horizon % self.model.o != 0:
|
if fc.max_horizon % self.model.o != 0:
|
||||||
@@ -505,8 +507,7 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
|
|||||||
"When compiling, max horizon needs to be multiple of the output patch"
|
"When compiling, max horizon needs to be multiple of the output patch"
|
||||||
" size %d. Using max horizon = %d instead.",
|
" size %d. Using max horizon = %d instead.",
|
||||||
self.model.o,
|
self.model.o,
|
||||||
new_horizon := math.ceil(fc.max_horizon / self.model.o)
|
new_horizon := math.ceil(fc.max_horizon / self.model.o) * self.model.o,
|
||||||
* self.model.o,
|
|
||||||
)
|
)
|
||||||
forecast_config = dataclasses.replace(fc, max_horizon=new_horizon)
|
forecast_config = dataclasses.replace(fc, max_horizon=new_horizon)
|
||||||
if fc.max_context + fc.max_horizon > self.model.config.context_limit:
|
if fc.max_context + fc.max_horizon > self.model.config.context_limit:
|
||||||
@@ -517,8 +518,7 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
|
|||||||
)
|
)
|
||||||
if fc.use_continuous_quantile_head and (fc.max_horizon > self.model.os):
|
if fc.use_continuous_quantile_head and (fc.max_horizon > self.model.os):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Continuous quantile head is not supported for horizons >"
|
f"Continuous quantile head is not supported for horizons > {self.model.os}."
|
||||||
f" {self.model.os}."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
self.forecast_config = forecast_config
|
self.forecast_config = forecast_config
|
||||||
@@ -534,19 +534,15 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def compiled_decode_kernel(fc, horizon, inputs, masks):
|
def compiled_decode_kernel(fc, horizon, inputs, masks):
|
||||||
|
|
||||||
inputs = jnp.array(inputs, dtype=jnp.float32)
|
inputs = jnp.array(inputs, dtype=jnp.float32)
|
||||||
masks = jnp.array(masks, dtype=jnp.bool)
|
masks = jnp.array(masks, dtype=jnp.bool)
|
||||||
if horizon > fc.max_horizon:
|
if horizon > fc.max_horizon:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Horizon must be less than the max horizon."
|
f"Horizon must be less than the max horizon. {horizon} > {fc.max_horizon}."
|
||||||
f" {horizon} > {fc.max_horizon}."
|
|
||||||
)
|
)
|
||||||
to_trim = fc.max_horizon - horizon
|
to_trim = fc.max_horizon - horizon
|
||||||
|
|
||||||
inputs, masks, is_positive, mu, sigma = _before_model_decode(
|
inputs, masks, is_positive, mu, sigma = _before_model_decode(fc, inputs, masks)
|
||||||
fc, inputs, masks
|
|
||||||
)
|
|
||||||
|
|
||||||
pf_outputs, quantile_spreads, ar_outputs = self.model.compiled_decode(
|
pf_outputs, quantile_spreads, ar_outputs = self.model.compiled_decode(
|
||||||
fc.max_horizon, inputs, masks
|
fc.max_horizon, inputs, masks
|
||||||
|
|||||||
@@ -55,12 +55,10 @@ class TimesFM_2p5_200M_torch_module(nn.Module):
|
|||||||
|
|
||||||
# Layers.
|
# Layers.
|
||||||
self.tokenizer = dense.ResidualBlock(self.config.tokenizer)
|
self.tokenizer = dense.ResidualBlock(self.config.tokenizer)
|
||||||
self.stacked_xf = nn.ModuleList(
|
self.stacked_xf = nn.ModuleList([
|
||||||
[
|
|
||||||
transformer.Transformer(self.config.stacked_transformers.transformer)
|
transformer.Transformer(self.config.stacked_transformers.transformer)
|
||||||
for _ in range(self.x)
|
for _ in range(self.x)
|
||||||
]
|
])
|
||||||
)
|
|
||||||
self.output_projection_point = dense.ResidualBlock(
|
self.output_projection_point = dense.ResidualBlock(
|
||||||
self.config.output_projection_point
|
self.config.output_projection_point
|
||||||
)
|
)
|
||||||
@@ -272,7 +270,7 @@ class TimesFM_2p5_200M_torch(timesfm_2p5_base.TimesFM_2p5, ModelHubMixin):
|
|||||||
def _from_pretrained(
|
def _from_pretrained(
|
||||||
cls,
|
cls,
|
||||||
*,
|
*,
|
||||||
model_id: str,
|
model_id: str = "google/timesfm-2.5-200m-pytorch",
|
||||||
revision: Optional[str],
|
revision: Optional[str],
|
||||||
cache_dir: Optional[Union[str, Path]],
|
cache_dir: Optional[Union[str, Path]],
|
||||||
force_download: bool,
|
force_download: bool,
|
||||||
|
|||||||
Reference in New Issue
Block a user