Add troubleshooting.md file and revert yapf changes

This commit is contained in:
Funto-Adeyemi
2023-11-12 21:38:11 +00:00
parent 1f9cb2bf92
commit 63f0bf4162
14 changed files with 979 additions and 919 deletions
+14 -12
View File
@@ -34,8 +34,9 @@ def get_seasonality(freq: str) -> int:
return _get_seasonality(freq, seasonalities={"D": 7})
def maybe_convert_col_to_datetime(df: pd.DataFrame,
col_name: str) -> pd.DataFrame:
def maybe_convert_col_to_datetime(
df: pd.DataFrame, col_name: str
) -> pd.DataFrame:
if not pd.api.types.is_datetime64_any_dtype(df[col_name]):
df = df.copy()
df[col_name] = pd.to_datetime(df[col_name])
@@ -63,15 +64,14 @@ def zero_pad_time_series(df, freq, min_length=36):
end=start_date,
periods=min_length - len(subset) + 1,
freq=freq, # 'MS' for month start
)[:-1] # Exclude the start_date itself
)[
:-1
] # Exclude the start_date itself
# 2c. Create padding data
padding_df = pd.DataFrame({
"ds": padding_dates,
"unique_id": unique_id,
"y": 0
} # Zero padding
)
padding_df = pd.DataFrame(
{"ds": padding_dates, "unique_id": unique_id, "y": 0} # Zero padding
)
# 2d. Combine original and padding data, and append to the list
padded_data.append(pd.concat([padding_df, subset]).sort_values("ds"))
@@ -121,7 +121,8 @@ class Forecaster:
for _, (cutoffs, train, valid) in tqdm(enumerate(splits)):
if len(valid.columns) > 3:
raise NotImplementedError(
"Cross validation with exogenous variables is not yet supported.")
"Cross validation with exogenous variables is not yet supported."
)
y_pred = self.forecast(
df=train,
h=h,
@@ -137,7 +138,8 @@ class Forecaster:
raise ValueError(
"Cross validation result produced less results than expected."
" Please verify that the frequency parameter (freq) matches your"
" series' and that there aren't any missing periods.")
" series' and that there aren't any missing periods."
)
results.append(result)
out = vertical_concat(results)
out = drop_index_if_pandas(out)
@@ -201,7 +203,7 @@ class TimeGPT(Forecaster):
all_unique_ids = df["unique_id"].unique()
all_fcst_df = []
for i in range(0, len(all_unique_ids), chunk_size):
chunk_ids = all_unique_ids[i:i + chunk_size]
chunk_ids = all_unique_ids[i : i + chunk_size]
chunk_df = df[df["unique_id"].isin(chunk_ids)]
fct_chunk_df = client.forecast(
df=chunk_df,
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Evaluation script for timegpt."""
import os
@@ -24,6 +25,7 @@ import pandas as pd
from ..baselines.timegpt_pipeline import run_timegpt
from .utils import ExperimentHandler
dataset_names = [
"m1_monthly",
"m1_quarterly",
@@ -61,6 +63,7 @@ _MODEL_NAME = flags.DEFINE_string(
)
_SAVE_DIR = flags.DEFINE_string("save_dir", "./results", "Save directory")
QUANTILES = list(np.arange(1, 10) / 10.0)
@@ -87,9 +90,9 @@ def main():
)
time_df = pd.DataFrame({"time": [total_time], "model": model_name})
fcsts_df = exp.fcst_from_level_to_quantiles(fcsts_df, model_name)
results = exp.evaluate_from_predictions(models=[model_name],
fcsts_df=fcsts_df,
times_df=time_df)
results = exp.evaluate_from_predictions(
models=[model_name], fcsts_df=fcsts_df, times_df=time_df
)
print(results, flush=True)
results_list.append(results)
results_full = pd.concat(results_list)
@@ -54,6 +54,7 @@ dataset_names = [
"hospital",
]
context_dict_v2 = {}
context_dict_v1 = {
+36 -24
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Forked from https://github.com/Nixtla/nixtla/blob/main/experiments/amazon-chronos/src/utils.py."""
from functools import partial
@@ -45,9 +46,11 @@ def quantile_loss(
target_col: str = "y",
) -> pd.DataFrame:
delta_y = df[models].sub(df[target_col], axis=0)
res = (np.maximum(q * delta_y,
(q - 1) * delta_y).groupby(df[id_col],
observed=True).mean())
res = (
np.maximum(q * delta_y, (q - 1) * delta_y)
.groupby(df[id_col], observed=True)
.mean()
)
res.index.name = id_col
res = res.reset_index()
return res
@@ -63,8 +66,10 @@ class ExperimentHandler:
models_dir: str = "./models",
):
if dataset not in gluonts_datasets:
raise Exception(f"dataset {dataset} not found in gluonts "
f"available datasets: {', '.join(gluonts_datasets)}")
raise Exception(
f"dataset {dataset} not found in gluonts "
f"available datasets: {', '.join(gluonts_datasets)}"
)
self.dataset = dataset
self.quantiles = quantiles
self.level = self._transform_quantiles_to_levels(quantiles)
@@ -75,8 +80,10 @@ class ExperimentHandler:
gluonts_dataset = get_dataset(self.dataset)
self.horizon = gluonts_dataset.metadata.prediction_length
if self.horizon is None:
raise Exception(f"horizon not found for dataset {self.dataset} "
"experiment cannot be run")
raise Exception(
f"horizon not found for dataset {self.dataset} "
"experiment cannot be run"
)
self.freq = gluonts_dataset.metadata.freq
# get_seasonality() returns 1 for freq='D', override this to 7. This significantly improves the accuracy of
# statistical models on datasets like m5/nn5_daily. The models like AutoARIMA/AutoETS can still set
@@ -115,8 +122,9 @@ class ExperimentHandler:
@staticmethod
def _transform_quantiles_to_levels(quantiles: List[float]) -> List[int]:
level = [int(100 - 200 * q) for q in quantiles if q < 0.5
] # in this case mean=mediain
level = [
int(100 - 200 * q) for q in quantiles if q < 0.5
] # in this case mean=mediain
level = sorted(list(set(level)))
return level
@@ -145,8 +153,9 @@ class ExperimentHandler:
last_n: int | None = None,
) -> pd.DataFrame:
with multiprocessing.Pool(os.cpu_count()) as pool: # Create a process pool
results = pool.map(parallel_transform, zip(gluonts_dataset,
repeat(last_n)))
results = pool.map(
parallel_transform, zip(gluonts_dataset, repeat(last_n))
)
df = pd.concat(results)
df = df.reset_index(drop=True)
return df
@@ -168,8 +177,9 @@ class ExperimentHandler:
def save_dataframe(self, df: pd.DataFrame, file_name: str):
df.to_csv(f"{self.results_dir}/{file_name}", index=False)
def save_results(self, fcst_df: pd.DataFrame, total_time: float,
model_name: str):
def save_results(
self, fcst_df: pd.DataFrame, total_time: float, model_name: str
):
self.save_dataframe(
fcst_df,
f"{model_name}-{self.dataset}-fcst.csv",
@@ -205,21 +215,23 @@ class ExperimentHandler:
times_df = []
for model in models:
fcst_method_df = pd.read_csv(
f"{self.results_dir}/{model}-{self.dataset}-fcst.csv").set_index(
["unique_id", "ds"])
f"{self.results_dir}/{model}-{self.dataset}-fcst.csv"
).set_index(["unique_id", "ds"])
fcsts_df.append(fcst_method_df)
time_method_df = pd.read_csv(
f"{self.results_dir}/{model}-{self.dataset}-time.csv")
f"{self.results_dir}/{model}-{self.dataset}-time.csv"
)
times_df.append(time_method_df)
fcsts_df = pd.concat(fcsts_df, axis=1).reset_index()
fcsts_df["ds"] = pd.to_datetime(fcsts_df["ds"])
times_df = pd.concat(times_df)
return self.evaluate_from_predictions(models=models,
fcsts_df=fcsts_df,
times_df=times_df)
return self.evaluate_from_predictions(
models=models, fcsts_df=fcsts_df, times_df=times_df
)
def evaluate_from_predictions(self, models: List[str], fcsts_df: pd.DataFrame,
times_df: pd.DataFrame) -> pd.DataFrame:
def evaluate_from_predictions(
self, models: List[str], fcsts_df: pd.DataFrame, times_df: pd.DataFrame
) -> pd.DataFrame:
test_df = self.test_df
train_df = self.train_df
test_df = test_df.merge(fcsts_df, how="left")
@@ -250,9 +262,9 @@ class ExperimentHandler:
eval_prob_df["metric"] = "scaled_crps"
eval_df = pd.concat([eval_df, eval_prob_df]).reset_index(drop=True)
eval_df = eval_df.groupby("metric").mean(numeric_only=True).reset_index()
eval_df = eval_df.melt(id_vars="metric",
value_name="value",
var_name="model")
eval_df = eval_df.melt(
id_vars="metric", value_name="value", var_name="model"
)
times_df.insert(0, "metric", "time")
times_df = times_df.rename(columns={"time": "value"})
eval_df = pd.concat([eval_df, times_df])
+286 -283
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
Finetune pipeline.
"""
@@ -38,11 +39,13 @@ from timesfm import TimesFm, data_loader, patched_decoder
NestedMap = py_utils.NestedMap
warnings.filterwarnings("ignore")
cmdstanpy_logger = logging.getLogger("cmdstanpy")
absl_logger = logging.getLogger("absl")
cmdstanpy_logger.disabled = True
absl_logger.disabled = True
"""
TimesFM model config. These are fixed since pre-training was done
with this configuration.
@@ -59,24 +62,20 @@ RANDOM_SEED = 1234
def finetune(
*,
model_name: Annotated[str,
typer.Option(
help="Specify the name of the huggingface model."
)] = "google/timesfm-1.0-200m",
model_name: Annotated[
str, typer.Option(help="Specify the name of the huggingface model.")
] = "google/timesfm-1.0-200m",
checkpoint_path: Annotated[
str,
typer.Option(help="The path to the local model checkpoint.")] = None,
datetime_col: Annotated[str,
typer.Option(
help="Column having datetime.")] = "ds",
ts_cols: Annotated[list[str],
typer.Option(
help="Columns of time-series features.")] = [],
normalize: Annotated[bool,
typer.Option(
help="Normalize data for eval or not")] = True,
context_len: Annotated[int,
typer.Option(help="Length of the context window")],
str, typer.Option(help="The path to the local model checkpoint.")
] = None,
datetime_col: Annotated[str, typer.Option(help="Column having datetime.")] = "ds",
ts_cols: Annotated[
list[str], typer.Option(help="Columns of time-series features.")
] = [],
normalize: Annotated[
bool, typer.Option(help="Normalize data for eval or not")
] = True,
context_len: Annotated[int, typer.Option(help="Length of the context window")],
horizon_len: Annotated[int, typer.Option(help="Prediction length.")],
freq: Annotated[
str,
@@ -88,312 +87,316 @@ def finetune(
data_path: Annotated[str, typer.Option(help="Path to dataset csv")],
boundaries: Annotated[
Tuple[int, int, int],
typer.Option(help="boundaries of dataset to train, val, test",),
typer.Option(
help="boundaries of dataset to train, val, test",
),
] = (0, 0, 0),
backend: Annotated[str,
typer.Option(help="Backend device: cpu, gpu, tpu")],
backend: Annotated[str, typer.Option(help="Backend device: cpu, gpu, tpu")],
batch_size: Annotated[
int,
typer.Option(help="Batch size for the randomly sampled batch")] = 16,
int, typer.Option(help="Batch size for the randomly sampled batch")
] = 16,
num_epochs: Annotated[int, typer.Option(help="Number of epochs")],
learning_rate: Annotated[float,
typer.Option(help="adam optimizer learning rate")],
adam_epsilon: Annotated[float,
typer.Option(help="adam optimizer epsilon")],
adam_clip_threshold: Annotated[float,
typer.Option(
help="adam optimizer clip threshold")],
cos_initial_decay_value: Annotated[float,
typer.Option(
help="cosine initial decay value")],
cos_final_decay_value: Annotated[float,
typer.Option(
help="cosine final decay value")],
cos_decay_steps: Annotated[int,
typer.Option(
help="Number of cosine decay steps")],
ema_decay: Annotated[float,
typer.Option(help="Exponential moving average decay")],
learning_rate: Annotated[float, typer.Option(help="adam optimizer learning rate")],
adam_epsilon: Annotated[float, typer.Option(help="adam optimizer epsilon")],
adam_clip_threshold: Annotated[
float, typer.Option(help="adam optimizer clip threshold")
],
cos_initial_decay_value: Annotated[
float, typer.Option(help="cosine initial decay value")
],
cos_final_decay_value: Annotated[
float, typer.Option(help="cosine final decay value")
],
cos_decay_steps: Annotated[int, typer.Option(help="Number of cosine decay steps")],
ema_decay: Annotated[float, typer.Option(help="Exponential moving average decay")],
early_stop_patience: Annotated[
int, typer.Option(..., help="Early stopping patience")] = 5,
int, typer.Option(..., help="Early stopping patience")
] = 5,
use_lora: Annotated[
bool,
typer.
Option(help="Train low rank adapters for stacked transformer block",),
typer.Option(
help="Train low rank adapters for stacked transformer block",
),
] = False,
lora_rank: Annotated[
int,
typer.Option(help="LoRA Rank",),
typer.Option(
help="LoRA Rank",
),
] = 8,
lora_target_modules: Annotated[
str,
typer.Option(
help=
"LoRA target modules of the transformer block. Allowed values: [all, attention, mlp]"
help="LoRA target modules of the transformer block. Allowed values: [all, attention, mlp]"
),
] = "all",
use_dora: Annotated[
bool,
typer.Option(help="Apply DoRA strategy along with LoRA.",),
typer.Option(
help="Apply DoRA strategy along with LoRA.",
),
] = False,
use_linear_probing: Annotated[
bool,
typer.Option(
help=
"Linear Probing. Train only input/output and embedding params. Freeze params in stack transformer block.",
help="Linear Probing. Train only input/output and embedding params. Freeze params in stack transformer block.",
),
] = False,
checkpoint_dir: Annotated[
str, typer.Option(help="Checkpoint directory")] = "./checkpoints",
wandb_project: Annotated[str,
typer.Option(help="Weights & Biases project name"
)] = "google_timesfm_finetune",
str, typer.Option(help="Checkpoint directory")
] = "./checkpoints",
wandb_project: Annotated[
str, typer.Option(help="Weights & Biases project name")
] = "google_timesfm_finetune",
) -> None:
key = jax.random.PRNGKey(seed=RANDOM_SEED)
wandb.init(project=wandb_project, config=locals())
key = jax.random.PRNGKey(seed=RANDOM_SEED)
wandb.init(project=wandb_project, config=locals())
data_df = pd.read_csv(open(data_path, "r"))
data_df = pd.read_csv(open(data_path, "r"))
if boundaries == (0, 0, 0):
# Default boundaries: train 60%, val 20%, test 20%
boundaries = [
int(len(data_df) * 0.6),
int(len(data_df) * 0.8),
len(data_df) - 1,
]
if boundaries == (0, 0, 0):
# Default boundaries: train 60%, val 20%, test 20%
boundaries = [
int(len(data_df) * 0.6),
int(len(data_df) * 0.8),
len(data_df) - 1,
]
ts_cols = [col for col in data_df.columns if col != datetime_col]
ts_cols = [col for col in data_df.columns if col != datetime_col]
dtl = data_loader.TimeSeriesdata(
data_path=data_path,
datetime_col=datetime_col,
num_cov_cols=None,
cat_cov_cols=None,
ts_cols=np.array(ts_cols),
train_range=[0, boundaries[0]],
val_range=[boundaries[0], boundaries[1]],
test_range=[boundaries[1], boundaries[2]],
hist_len=context_len,
pred_len=horizon_len,
batch_size=batch_size,
freq=freq,
normalize=normalize,
epoch_len=None,
holiday=False,
permute=False,
)
train_batches = dtl.tf_dataset(mode="train", shift=1).batch(batch_size)
val_batches = dtl.tf_dataset(mode="val", shift=horizon_len)
for tbatch in tqdm(train_batches.as_numpy_iterator()):
pass
tfm = TimesFm(
context_len=context_len,
horizon_len=horizon_len,
input_patch_len=INPUT_PATCH_LEN,
output_patch_len=OUTPUT_PATCH_LEN,
num_layers=NUM_LAYERS,
model_dims=MODEL_DIMS,
backend=backend,
per_core_batch_size=batch_size,
quantiles=QUANTILES,
)
if checkpoint_path:
tfm.load_from_checkpoint(
checkpoint_path=checkpoint_path,
checkpoint_type=checkpoints.CheckpointType.FLAX,
)
else:
tfm.load_from_checkpoint(
repo_id=model_name,
checkpoint_type=checkpoints.CheckpointType.FLAX,
dtl = data_loader.TimeSeriesdata(
data_path=data_path,
datetime_col=datetime_col,
num_cov_cols=None,
cat_cov_cols=None,
ts_cols=np.array(ts_cols),
train_range=[0, boundaries[0]],
val_range=[boundaries[0], boundaries[1]],
test_range=[boundaries[1], boundaries[2]],
hist_len=context_len,
pred_len=horizon_len,
batch_size=batch_size,
freq=freq,
normalize=normalize,
epoch_len=None,
holiday=False,
permute=False,
)
model = pax_fiddle.Config(
patched_decoder.PatchedDecoderFinetuneModel,
name="patched_decoder_finetune",
core_layer_tpl=tfm.model_p,
)
train_batches = dtl.tf_dataset(mode="train", shift=1).batch(batch_size)
val_batches = dtl.tf_dataset(mode="val", shift=horizon_len)
if use_lora:
load_adapter_layer(
mdl_vars=tfm._train_state.mdl_vars,
model=model.core_layer_tpl,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
use_dora=use_dora,
for tbatch in tqdm(train_batches.as_numpy_iterator()):
pass
tfm = TimesFm(
context_len=context_len,
horizon_len=horizon_len,
input_patch_len=INPUT_PATCH_LEN,
output_patch_len=OUTPUT_PATCH_LEN,
num_layers=NUM_LAYERS,
model_dims=MODEL_DIMS,
backend=backend,
per_core_batch_size=batch_size,
quantiles=QUANTILES,
)
if checkpoint_path:
tfm.load_from_checkpoint(
checkpoint_path=checkpoint_path,
checkpoint_type=checkpoints.CheckpointType.FLAX,
)
else:
tfm.load_from_checkpoint(
repo_id=model_name,
checkpoint_type=checkpoints.CheckpointType.FLAX,
)
model = pax_fiddle.Config(
patched_decoder.PatchedDecoderFinetuneModel,
name="patched_decoder_finetune",
core_layer_tpl=tfm.model_p,
)
@pax_fiddle.auto_config
def build_learner() -> learners.Learner:
bprop_variable_inclusion = []
bprop_variable_exclusion = []
if use_lora:
bprop_variable_inclusion.append(r"^.*lora.*$")
if use_dora:
bprop_variable_inclusion.append(r"^.*dora.*$")
elif use_linear_probing:
bprop_variable_exclusion = [".*/stacked_transformer_layer/.*"]
return pax_fiddle.Config(
learners.Learner,
name="learner",
loss_name="avg_qloss",
optimizer=optimizers.Adam(
epsilon=adam_epsilon,
clip_threshold=adam_clip_threshold,
learning_rate=learning_rate,
lr_schedule=pax_fiddle.Config(
schedules.Cosine,
initial_value=cos_initial_decay_value,
final_value=cos_final_decay_value,
total_steps=cos_decay_steps,
),
ema_decay=ema_decay,
),
bprop_variable_exclusion=bprop_variable_exclusion,
bprop_variable_inclusion=bprop_variable_inclusion,
)
task_p = tasks_lib.SingleTask(
name="ts-learn",
model=model,
train=tasks_lib.SingleTask.Train(learner=build_learner(),),
)
task_p.model.ici_mesh_shape = [1, 1, 1]
task_p.model.mesh_axis_names = ["replica", "data", "mdl"]
DEVICES = np.array(jax.devices()).reshape([1, 1, 1])
jax.sharding.Mesh(DEVICES, ["replica", "data", "mdl"])
num_devices = jax.local_device_count()
print(f"num_devices: {num_devices}")
print(f"device kind: {jax.local_devices()[0].device_kind}")
jax_task = task_p
key, init_key = jax.random.split(key)
def process_train_batch(batch):
past_ts = batch[0].reshape(batch_size * len(ts_cols), -1)
actual_ts = batch[3].reshape(batch_size * len(ts_cols), -1)
return NestedMap(input_ts=past_ts, actual_ts=actual_ts)
def process_eval_batch(batch):
past_ts = batch[0]
actual_ts = batch[3]
return NestedMap(input_ts=past_ts, actual_ts=actual_ts)
jax_model_states, _ = trainer_lib.initialize_model_state(
jax_task,
init_key,
process_train_batch(tbatch),
checkpoint_type=checkpoint_types.CheckpointType.GDA,
)
jax_model_states.mdl_vars["params"]["core_layer"] = tfm._train_state.mdl_vars[
"params"]
gc.collect()
jax_task = task_p
def train_step(states, prng_key, inputs):
return trainer_lib.train_step_single_learner(jax_task, states, prng_key,
inputs)
def eval_step(states, prng_key, inputs):
states = states.to_eval_state()
return trainer_lib.eval_step_single_learner(jax_task, states, prng_key,
inputs)
key, train_key, eval_key = jax.random.split(key, 3)
train_prng_seed = jax.random.split(train_key, num=jax.local_device_count())
eval_prng_seed = jax.random.split(eval_key, num=jax.local_device_count())
p_train_step = jax.pmap(train_step, axis_name="batch")
p_eval_step = jax.pmap(eval_step, axis_name="batch")
replicated_jax_states = trainer_lib.replicate_model_state(jax_model_states)
def reshape_batch_for_pmap(batch, num_devices):
def _reshape(input_tensor):
bsize = input_tensor.shape[0]
residual_shape = list(input_tensor.shape[1:])
nbsize = bsize // num_devices
return jnp.reshape(input_tensor, [num_devices, nbsize] + residual_shape)
return jax.tree.map(_reshape, batch)
patience = 0
best_eval_loss = 1e7
checkpoint_dir = f"{checkpoint_dir}/run_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{wandb.run.id}"
for epoch in range(num_epochs):
if patience >= early_stop_patience:
print("Early stopping.")
break
print(f"Epoch: {epoch + 1}")
train_its = train_batches.as_numpy_iterator()
train_losses = []
for batch in tqdm(train_its):
tbatch = process_train_batch(batch)
tbatch = reshape_batch_for_pmap(tbatch, num_devices)
replicated_jax_states, step_fun_out = p_train_step(
replicated_jax_states, train_prng_seed, tbatch)
train_losses.append(step_fun_out.loss[0])
wandb.log({"train_step_loss": step_fun_out.loss[0]})
avg_train_loss = np.mean(train_losses)
print("Starting eval.")
val_its = val_batches.as_numpy_iterator()
eval_losses = []
for ev_batch in tqdm(val_its):
ebatch = process_eval_batch(ev_batch)
ebatch = reshape_batch_for_pmap(ebatch, num_devices)
_, step_fun_out = p_eval_step(replicated_jax_states, eval_prng_seed,
ebatch)
eval_losses.append(step_fun_out.loss[0])
wandb.log({"eval_step_loss": step_fun_out.loss[0]})
avg_eval_loss = np.mean(eval_losses)
print(f"Train Loss: {avg_train_loss}, Val Loss: {avg_eval_loss}")
wandb.log({
"epoch": epoch + 1,
"avg_train_loss": avg_train_loss,
"avg_val_loss": avg_eval_loss,
})
if avg_eval_loss < best_eval_loss or np.isnan(avg_eval_loss):
best_eval_loss = avg_eval_loss
print("Saving checkpoint.")
jax_state_for_saving = py_utils.maybe_unreplicate_for_fully_replicated(
replicated_jax_states)
if use_lora:
adapter_params = get_adapter_params(
params=jax_state_for_saving.mdl_vars,
load_adapter_layer(
mdl_vars=tfm._train_state.mdl_vars,
model=model.core_layer_tpl,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
num_layers=NUM_LAYERS,
use_dora=use_dora,
)
jax_state_for_saving.mdl_vars["params"] = adapter_params
checkpoints.save_checkpoint(jax_state_for_saving,
checkpoint_dir,
overwrite=True)
@pax_fiddle.auto_config
def build_learner() -> learners.Learner:
bprop_variable_inclusion = []
bprop_variable_exclusion = []
if use_lora:
bprop_variable_inclusion.append(r"^.*lora.*$")
if use_dora:
bprop_variable_inclusion.append(r"^.*dora.*$")
elif use_linear_probing:
bprop_variable_exclusion = [".*/stacked_transformer_layer/.*"]
patience = 0
del jax_state_for_saving
gc.collect()
else:
patience += 1
print(f"patience: {patience}")
print("Fine-tuning completed.")
return pax_fiddle.Config(
learners.Learner,
name="learner",
loss_name="avg_qloss",
optimizer=optimizers.Adam(
epsilon=adam_epsilon,
clip_threshold=adam_clip_threshold,
learning_rate=learning_rate,
lr_schedule=pax_fiddle.Config(
schedules.Cosine,
initial_value=cos_initial_decay_value,
final_value=cos_final_decay_value,
total_steps=cos_decay_steps,
),
ema_decay=ema_decay,
),
bprop_variable_exclusion=bprop_variable_exclusion,
bprop_variable_inclusion=bprop_variable_inclusion,
)
task_p = tasks_lib.SingleTask(
name="ts-learn",
model=model,
train=tasks_lib.SingleTask.Train(
learner=build_learner(),
),
)
task_p.model.ici_mesh_shape = [1, 1, 1]
task_p.model.mesh_axis_names = ["replica", "data", "mdl"]
DEVICES = np.array(jax.devices()).reshape([1, 1, 1])
jax.sharding.Mesh(DEVICES, ["replica", "data", "mdl"])
num_devices = jax.local_device_count()
print(f"num_devices: {num_devices}")
print(f"device kind: {jax.local_devices()[0].device_kind}")
jax_task = task_p
key, init_key = jax.random.split(key)
def process_train_batch(batch):
past_ts = batch[0].reshape(batch_size * len(ts_cols), -1)
actual_ts = batch[3].reshape(batch_size * len(ts_cols), -1)
return NestedMap(input_ts=past_ts, actual_ts=actual_ts)
def process_eval_batch(batch):
past_ts = batch[0]
actual_ts = batch[3]
return NestedMap(input_ts=past_ts, actual_ts=actual_ts)
jax_model_states, _ = trainer_lib.initialize_model_state(
jax_task,
init_key,
process_train_batch(tbatch),
checkpoint_type=checkpoint_types.CheckpointType.GDA,
)
jax_model_states.mdl_vars["params"]["core_layer"] = tfm._train_state.mdl_vars[
"params"
]
gc.collect()
jax_task = task_p
def train_step(states, prng_key, inputs):
return trainer_lib.train_step_single_learner(jax_task, states, prng_key, inputs)
def eval_step(states, prng_key, inputs):
states = states.to_eval_state()
return trainer_lib.eval_step_single_learner(jax_task, states, prng_key, inputs)
key, train_key, eval_key = jax.random.split(key, 3)
train_prng_seed = jax.random.split(train_key, num=jax.local_device_count())
eval_prng_seed = jax.random.split(eval_key, num=jax.local_device_count())
p_train_step = jax.pmap(train_step, axis_name="batch")
p_eval_step = jax.pmap(eval_step, axis_name="batch")
replicated_jax_states = trainer_lib.replicate_model_state(jax_model_states)
def reshape_batch_for_pmap(batch, num_devices):
def _reshape(input_tensor):
bsize = input_tensor.shape[0]
residual_shape = list(input_tensor.shape[1:])
nbsize = bsize // num_devices
return jnp.reshape(input_tensor, [num_devices, nbsize] + residual_shape)
return jax.tree.map(_reshape, batch)
patience = 0
best_eval_loss = 1e7
checkpoint_dir = f"{checkpoint_dir}/run_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{wandb.run.id}"
for epoch in range(num_epochs):
if patience >= early_stop_patience:
print("Early stopping.")
break
print(f"Epoch: {epoch + 1}")
train_its = train_batches.as_numpy_iterator()
train_losses = []
for batch in tqdm(train_its):
tbatch = process_train_batch(batch)
tbatch = reshape_batch_for_pmap(tbatch, num_devices)
replicated_jax_states, step_fun_out = p_train_step(
replicated_jax_states, train_prng_seed, tbatch
)
train_losses.append(step_fun_out.loss[0])
wandb.log({"train_step_loss": step_fun_out.loss[0]})
avg_train_loss = np.mean(train_losses)
print("Starting eval.")
val_its = val_batches.as_numpy_iterator()
eval_losses = []
for ev_batch in tqdm(val_its):
ebatch = process_eval_batch(ev_batch)
ebatch = reshape_batch_for_pmap(ebatch, num_devices)
_, step_fun_out = p_eval_step(replicated_jax_states, eval_prng_seed, ebatch)
eval_losses.append(step_fun_out.loss[0])
wandb.log({"eval_step_loss": step_fun_out.loss[0]})
avg_eval_loss = np.mean(eval_losses)
print(f"Train Loss: {avg_train_loss}, Val Loss: {avg_eval_loss}")
wandb.log(
{
"epoch": epoch + 1,
"avg_train_loss": avg_train_loss,
"avg_val_loss": avg_eval_loss,
}
)
if avg_eval_loss < best_eval_loss or np.isnan(avg_eval_loss):
best_eval_loss = avg_eval_loss
print("Saving checkpoint.")
jax_state_for_saving = py_utils.maybe_unreplicate_for_fully_replicated(
replicated_jax_states
)
if use_lora:
adapter_params = get_adapter_params(
params=jax_state_for_saving.mdl_vars,
lora_target_modules=lora_target_modules,
num_layers=NUM_LAYERS,
use_dora=use_dora,
)
jax_state_for_saving.mdl_vars["params"] = adapter_params
checkpoints.save_checkpoint(
jax_state_for_saving, checkpoint_dir, overwrite=True
)
patience = 0
del jax_state_for_saving
gc.collect()
else:
patience += 1
print(f"patience: {patience}")
print("Fine-tuning completed.")
if __name__ == "__main__":
typer.run(finetune)
typer.run(finetune)
+1
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""adapter init file."""
from .dora_layers import DoraAttentionProjection, DoraCombinedQKVProjection, DoraLinear
+149 -148
View File
@@ -21,181 +21,182 @@ WeightHParams = base_layer.WeightHParams
class DoraTheta(base_layer.Theta):
def __init__(self, module):
self.module = module
def __init__(self, module):
self.module = module
def _dora_initialized(self):
if (
self.module.has_variable("params", "lora_a")
and self.module.has_variable("params", "lora_b")
and self.module.has_variable("params", "dora_m")
and "lora_a" in self.module._weight_hparams
and "lora_b" in self.module._weight_hparams
and "dora_m" in self.module._weight_hparams
):
return True
else:
return False
def _dora_initialized(self):
if (self.module.has_variable("params", "lora_a") and
self.module.has_variable("params", "lora_b") and
self.module.has_variable("params", "dora_m") and
"lora_a" in self.module._weight_hparams and
"lora_b" in self.module._weight_hparams and
"dora_m" in self.module._weight_hparams):
return True
else:
return False
def _dorafy_var(self, w):
lora_a = super().__getattr__("lora_a")
lora_b = super().__getattr__("lora_b")
dora_m = super().__getattr__("dora_m")
def _dorafy_var(self, w):
lora_a = super().__getattr__("lora_a")
lora_b = super().__getattr__("lora_b")
dora_m = super().__getattr__("dora_m")
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
w_prime = w + lora_delta
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
return w_prime
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
return w_prime
def __getattr__(self, k):
var = super().__getattr__(k)
if not self._dora_initialized():
return var
def __getattr__(self, k):
var = super().__getattr__(k)
if not self._dora_initialized():
return var
if k == "w":
return self._dorafy_var(var)
if k == "w":
return self._dorafy_var(var)
return var
return var
def __getitem__(self, k):
var = super().__getattr__(k)
if not self._dora_initialized():
return var
def __getitem__(self, k):
var = super().__getattr__(k)
if not self._dora_initialized():
return var
if k == "w":
return self._dorafy_var(var)
if k == "w":
return self._dorafy_var(var)
return var
return var
class DoraThetaDescriptor:
"""Dot syntax accession descriptor."""
"""Dot syntax accession descriptor."""
def __get__(self, obj, objtype=None):
return DoraTheta(obj)
def __get__(self, obj, objtype=None):
return DoraTheta(obj)
class DoraLinear(linears.Linear):
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init
def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init
super().setup()
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dims, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[1, self.output_dims],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
super().setup()
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dims, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[1, self.output_dims],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
class DoraAttentionProjection(attentions.AttentionProjection):
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[1, self.num_heads, self.dim_per_head],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[1, self.num_heads, self.dim_per_head],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
class DoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer):
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[3, self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[3, 1, self.num_heads, self.dim_per_head],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None, None],
),
)
self.create_variable(
"lora_a",
WeightHParams(
shape=[3, self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[3, 1, self.num_heads, self.dim_per_head],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None, None],
),
)
+116 -115
View File
@@ -21,145 +21,146 @@ WeightHParams = base_layer.WeightHParams
class LoraTheta(base_layer.Theta):
def __init__(self, module):
self.module = module
def __init__(self, module):
self.module = module
def _lora_initialized(self):
if (
self.module.has_variable("params", "lora_a")
and self.module.has_variable("params", "lora_b")
and "lora_a" in self.module._weight_hparams
and "lora_b" in self.module._weight_hparams
):
return True
else:
return False
def _lora_initialized(self):
if (self.module.has_variable("params", "lora_a") and
self.module.has_variable("params", "lora_b") and
"lora_a" in self.module._weight_hparams and
"lora_b" in self.module._weight_hparams):
return True
else:
return False
def _lorafy_var(self, w):
lora_a = super().__getattr__("lora_a")
lora_b = super().__getattr__("lora_b")
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
return w_prime
def _lorafy_var(self, w):
lora_a = super().__getattr__("lora_a")
lora_b = super().__getattr__("lora_b")
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
return w_prime
def __getattr__(self, k):
var = super().__getattr__(k)
if not self._lora_initialized():
return var
def __getattr__(self, k):
var = super().__getattr__(k)
if not self._lora_initialized():
return var
if k == "w":
return self._lorafy_var(var)
if k == "w":
return self._lorafy_var(var)
return var
return var
def __getitem__(self, k):
var = super().__getattr__(k)
if not self._lora_initialized():
return var
def __getitem__(self, k):
var = super().__getattr__(k)
if not self._lora_initialized():
return var
if k == "w":
return self._lorafy_var(var)
if k == "w":
return self._lorafy_var(var)
return var
return var
class LoraThetaDescriptor:
"""Dot syntax accession descriptor."""
"""Dot syntax accession descriptor."""
def __get__(self, obj, objtype=None):
return LoraTheta(obj)
def __get__(self, obj, objtype=None):
return LoraTheta(obj)
class LoraLinear(linears.Linear):
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init
def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init
super().setup()
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dims, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
super().setup()
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dims, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
class LoraAttentionProjection(attentions.AttentionProjection):
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
class LoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer):
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[3, self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_a",
WeightHParams(
shape=[3, self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
+286 -256
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
This file provides functionality for loading and merging adapter weights
in timesfm model, specifically for LoRA and DoRA.
@@ -39,11 +40,10 @@ from adapter.lora_layers import (
from timesfm import TimesFm
def get_adapter_params(params: dict,
lora_target_modules: str,
num_layers: int,
use_dora: bool = False) -> dict:
"""
def get_adapter_params(
params: dict, lora_target_modules: str, num_layers: int, use_dora: bool = False
) -> dict:
"""
Extracts adapter parameters from the given model parameters for saving the checkpoint.
Args:
@@ -55,44 +55,47 @@ def get_adapter_params(params: dict,
Returns:
dict: A dictionary containing the extracted adapter parameters.
"""
adapter_params = {}
for i in range(num_layers):
layer_key = f"x_layers_{i}"
adapter_params[layer_key] = {}
adapter_params = {}
for i in range(num_layers):
layer_key = f"x_layers_{i}"
adapter_params[layer_key] = {}
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key]["ff_layer"][ff_layer_key]["linear"]
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key
]["ff_layer"][ff_layer_key]["linear"]
lora_a = linear["lora_a"]
lora_b = linear["lora_b"]
lora_a = linear["lora_a"]
lora_b = linear["lora_b"]
adapter_params[layer_key][ff_layer_key] = {
"lora_a": lora_a,
"lora_b": lora_b,
}
adapter_params[layer_key][ff_layer_key] = {
"lora_a": lora_a,
"lora_b": lora_b,
}
if use_dora:
adapter_params[layer_key][ff_layer_key]["dora_m"] = linear["dora_m"]
if use_dora:
adapter_params[layer_key][ff_layer_key]["dora_m"] = linear["dora_m"]
if lora_target_modules in ["all", "attention"]:
attention = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key]["self_attention"]
if lora_target_modules in ["all", "attention"]:
attention = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key
]["self_attention"]
for component in ["key", "query", "value", "post"]:
lora_a = attention[component]["lora_a"]
lora_b = attention[component]["lora_b"]
for component in ["key", "query", "value", "post"]:
lora_a = attention[component]["lora_a"]
lora_b = attention[component]["lora_b"]
adapter_params[layer_key][component] = {
"lora_a": lora_a,
"lora_b": lora_b,
}
adapter_params[layer_key][component] = {
"lora_a": lora_a,
"lora_b": lora_b,
}
if use_dora:
adapter_params[layer_key][component]["dora_m"] = attention[component][
"dora_m"]
return adapter_params
if use_dora:
adapter_params[layer_key][component]["dora_m"] = attention[
component
]["dora_m"]
return adapter_params
def load_adapter_checkpoint(
@@ -102,7 +105,7 @@ def load_adapter_checkpoint(
lora_target_modules: str,
use_dora: bool,
) -> None:
"""
"""
Loads an adapter checkpoint and merges it with the original model weights.
Args:
@@ -115,77 +118,83 @@ def load_adapter_checkpoint(
Returns:
None
"""
"""
"""
currently loading and initializing the model with adapter layers first and then merging the
adapter weights to original weights and replacing the adapter layers back to original layer.
# NOTE: refactor this. there should be a better way to load the LoRA checkpoint.
"""
model._logging(
f"Restoring adapter checkpoint from {adapter_checkpoint_path}.")
start_time = time.time()
original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl = (
load_adapter_layer(
mdl_vars=model._train_state.mdl_vars,
model=model._model,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
use_dora=use_dora,
))
model._logging(f"Restoring adapter checkpoint from {adapter_checkpoint_path}.")
start_time = time.time()
original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl = (
load_adapter_layer(
mdl_vars=model._train_state.mdl_vars,
model=model._model,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
use_dora=use_dora,
)
)
var_weight_hparams = model._model.abstract_init_with_metadata(
model._get_sample_inputs(), do_eval=True)
var_weight_hparams = model._model.abstract_init_with_metadata(
model._get_sample_inputs(), do_eval=True
)
adapter_weight_hparams = _get_adapter_weight_params(
var_weight_hparams=var_weight_hparams,
lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora,
)
adapter_weight_hparams = _get_adapter_weight_params(
var_weight_hparams=var_weight_hparams,
lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora,
)
adapter_state_partition_specs = tasks_lib.create_state_partition_specs(
adapter_weight_hparams,
mesh_shape=model.mesh_shape,
mesh_axis_names=model.mesh_name,
discard_opt_states=True,
learners=None,
)
adapter_state_local_shapes = tasks_lib.create_state_unpadded_shapes(
adapter_weight_hparams,
discard_opt_states=True,
learners=None,
)
adapter_train_state = checkpoints.restore_checkpoint(
state_global_shapes=adapter_state_local_shapes,
checkpoint_dir=adapter_checkpoint_path,
checkpoint_type=checkpoints.CheckpointType.FLAX,
state_specs=adapter_state_partition_specs,
step=None,
)
adapter_state_partition_specs = tasks_lib.create_state_partition_specs(
adapter_weight_hparams,
mesh_shape=model.mesh_shape,
mesh_axis_names=model.mesh_name,
discard_opt_states=True,
learners=None,
)
adapter_state_local_shapes = tasks_lib.create_state_unpadded_shapes(
adapter_weight_hparams,
discard_opt_states=True,
learners=None,
)
adapter_train_state = checkpoints.restore_checkpoint(
state_global_shapes=adapter_state_local_shapes,
checkpoint_dir=adapter_checkpoint_path,
checkpoint_type=checkpoints.CheckpointType.FLAX,
state_specs=adapter_state_partition_specs,
step=None,
)
# add adapter weights to the original weights
_merge_adapter_weights(
model=model,
adapter_train_state=adapter_train_state,
lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora,
)
# add adapter weights to the original weights
_merge_adapter_weights(
model=model,
adapter_train_state=adapter_train_state,
lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora,
)
# replace back with the original model layer
if lora_target_modules in ["all", "mlp"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
original_linear_tpl)
# replace back with the original model layer
if lora_target_modules in ["all", "mlp"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
original_linear_tpl
)
if lora_target_modules in ["all", "attention"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
original_attn_tpl)
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
original_combined_qkv_tpl)
model._logging(
f"Restored adapter checkpoint in {time.time() - start_time:.2f} seconds.")
if lora_target_modules in ["all", "attention"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
original_attn_tpl
)
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
original_combined_qkv_tpl
)
model._logging(
f"Restored adapter checkpoint in {time.time() - start_time:.2f} seconds."
)
# jit compile the model
model.jit_decode()
# jit compile the model
model.jit_decode()
def _merge_adapter_weights(
@@ -195,7 +204,7 @@ def _merge_adapter_weights(
num_layers: int,
use_dora: bool,
) -> None:
"""
"""
Merges adapter weights with the original model weights.
Args:
@@ -205,73 +214,74 @@ def _merge_adapter_weights(
num_layers (int): Number of transformer layers.
use_dora (bool): Whether DoRA was used or not.
"""
for i in range(num_layers):
layer_key = f"x_layers_{i}"
for i in range(num_layers):
layer_key = f"x_layers_{i}"
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"][layer_key]["ff_layer"][ff_layer_key][
"linear"]
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"
][layer_key]["ff_layer"][ff_layer_key]["linear"]
params = adapter_train_state.mdl_vars[layer_key][ff_layer_key]
lora_a = params["lora_a"]
lora_b = params["lora_b"]
params = adapter_train_state.mdl_vars[layer_key][ff_layer_key]
lora_a = params["lora_a"]
lora_b = params["lora_b"]
w = linear["w"]
w = linear["w"]
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
if use_dora:
dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
linear["w"] = w_prime
del linear["dora_m"]
if use_dora:
dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
linear["w"] = w_prime
del linear["dora_m"]
else:
linear["w"] = w_prime
else:
linear["w"] = w_prime
del linear["lora_a"]
del linear["lora_b"]
del linear["lora_a"]
del linear["lora_b"]
if lora_target_modules in ["all", "attention"]:
attention = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"][layer_key]["self_attention"]
if lora_target_modules in ["all", "attention"]:
attention = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"
][layer_key]["self_attention"]
for component in ["key", "query", "value", "post"]:
params = adapter_train_state.mdl_vars[layer_key][component]
lora_a = params["lora_a"]
lora_b = params["lora_b"]
for component in ["key", "query", "value", "post"]:
params = adapter_train_state.mdl_vars[layer_key][component]
lora_a = params["lora_a"]
lora_b = params["lora_b"]
w = attention[component]["w"]
w = attention[component]["w"]
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
if use_dora:
dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
attention[component]["w"] = w_prime
del attention[component]["dora_m"]
if use_dora:
dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
attention[component]["w"] = w_prime
del attention[component]["dora_m"]
else:
attention[component]["w"] = w_prime
else:
attention[component]["w"] = w_prime
del attention[component]["lora_a"]
del attention[component]["lora_b"]
del attention[component]["lora_a"]
del attention[component]["lora_b"]
def _get_adapter_weight_params(var_weight_hparams: dict,
lora_target_modules: str, num_layers: int,
use_dora: bool) -> dict:
"""
def _get_adapter_weight_params(
var_weight_hparams: dict, lora_target_modules: str, num_layers: int, use_dora: bool
) -> dict:
"""
Extracts adapter weight parameters from the given variable weight hyperparameters.
Args:
@@ -283,39 +293,42 @@ def _get_adapter_weight_params(var_weight_hparams: dict,
Returns:
dict: A dictionary containing the extracted adapter weight parameters.
"""
adapter_params = {}
for i in range(num_layers):
layer = f"x_layers_{i}"
adapter_params[layer] = {}
adapter_params = {}
for i in range(num_layers):
layer = f"x_layers_{i}"
adapter_params[layer] = {}
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"][layer]["ff_layer"][ff_layer_key][
"linear"]
adapter_params[layer][ff_layer_key] = {
"lora_a": adapter_weight_params["lora_a"],
"lora_b": adapter_weight_params["lora_b"],
}
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"
][layer]["ff_layer"][ff_layer_key]["linear"]
adapter_params[layer][ff_layer_key] = {
"lora_a": adapter_weight_params["lora_a"],
"lora_b": adapter_weight_params["lora_b"],
}
if use_dora:
adapter_params[layer][ff_layer_key]["dora_m"] = (
adapter_weight_params["dora_m"])
if use_dora:
adapter_params[layer][ff_layer_key]["dora_m"] = (
adapter_weight_params["dora_m"]
)
if lora_target_modules in ["all", "attention"]:
for component in ["key", "value", "query", "post"]:
adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"][layer]["self_attention"][component]
adapter_params[layer][component] = {
"lora_a": adapter_weight_params["lora_a"],
"lora_b": adapter_weight_params["lora_b"],
}
if lora_target_modules in ["all", "attention"]:
for component in ["key", "value", "query", "post"]:
adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"
][layer]["self_attention"][component]
adapter_params[layer][component] = {
"lora_a": adapter_weight_params["lora_a"],
"lora_b": adapter_weight_params["lora_b"],
}
if use_dora:
adapter_params[layer][component]["dora_m"] = adapter_weight_params[
"dora_m"]
if use_dora:
adapter_params[layer][component]["dora_m"] = adapter_weight_params[
"dora_m"
]
return adapter_params
return adapter_params
def load_adapter_layer(
@@ -325,7 +338,7 @@ def load_adapter_layer(
lora_target_modules: str,
use_dora: bool = False,
) -> tuple[pax_fiddle.Config, pax_fiddle.Config]:
"""
"""
Updates target modules with adapter layers.
Args:
@@ -338,55 +351,67 @@ def load_adapter_layer(
Returns:
tuple[pax_fiddle.Config, pax_fiddle.Config]: Updated model configurations.
"""
original_linear_tpl = original_attn_tpl = original_combined_qkv_tpl = None
if lora_target_modules in ["all", "mlp"]:
original_linear_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.
tr_fflayer_tpl.fflayer_tpl.linear_tpl)
adapter_linear_tpl = (pax_fiddle.Config(
DoraLinear,
rank=lora_rank,
) if use_dora else pax_fiddle.Config(
LoraLinear,
rank=lora_rank,
))
adapter_linear_tpl.copy_fields_from(original_linear_tpl)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
adapter_linear_tpl)
original_linear_tpl = original_attn_tpl = original_combined_qkv_tpl = None
if lora_target_modules in ["all", "mlp"]:
original_linear_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl
)
adapter_linear_tpl = (
pax_fiddle.Config(
DoraLinear,
rank=lora_rank,
)
if use_dora
else pax_fiddle.Config(
LoraLinear,
rank=lora_rank,
)
)
adapter_linear_tpl.copy_fields_from(original_linear_tpl)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
adapter_linear_tpl
)
if lora_target_modules in ["all", "attention"]:
original_attn_tpl = (model.stacked_transformer_params_tpl.
transformer_layer_params_tpl.tr_atten_tpl.proj_tpl)
if lora_target_modules in ["all", "attention"]:
original_attn_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl
)
adapter_attn_tpl = (
pax_fiddle.Config(DoraAttentionProjection, rank=lora_rank) if use_dora
else pax_fiddle.Config(LoraAttentionProjection, rank=lora_rank))
adapter_attn_tpl.copy_fields_from(original_attn_tpl)
adapter_attn_tpl = (
pax_fiddle.Config(DoraAttentionProjection, rank=lora_rank)
if use_dora
else pax_fiddle.Config(LoraAttentionProjection, rank=lora_rank)
)
adapter_attn_tpl.copy_fields_from(original_attn_tpl)
original_combined_qkv_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.
tr_atten_tpl.combined_qkv_proj_tpl)
original_combined_qkv_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl
)
adapter_combined_qkv_tpl = (
pax_fiddle.Config(DoraCombinedQKVProjection, rank=lora_rank) if use_dora
else pax_fiddle.Config(LoraCombinedQKVProjection, rank=lora_rank))
adapter_combined_qkv_tpl.copy_fields_from(original_combined_qkv_tpl)
adapter_combined_qkv_tpl = (
pax_fiddle.Config(DoraCombinedQKVProjection, rank=lora_rank)
if use_dora
else pax_fiddle.Config(LoraCombinedQKVProjection, rank=lora_rank)
)
adapter_combined_qkv_tpl.copy_fields_from(original_combined_qkv_tpl)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
adapter_attn_tpl)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
adapter_combined_qkv_tpl)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
adapter_attn_tpl
)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
adapter_combined_qkv_tpl
)
# initialize and add adapter weights
_initialize_adapter_params(
mdl_vars=mdl_vars,
num_layers=model.stacked_transformer_params_tpl.num_layers,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
use_dora=use_dora,
)
# initialize and add adapter weights
_initialize_adapter_params(
mdl_vars=mdl_vars,
num_layers=model.stacked_transformer_params_tpl.num_layers,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
use_dora=use_dora,
)
return original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl
return original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl
def _initialize_adapter_params(
@@ -397,7 +422,7 @@ def _initialize_adapter_params(
use_dora: bool = False,
seed: int = 1234,
) -> dict:
"""
"""
Initializes and adds adapter parameters to target modules.
Args:
@@ -411,47 +436,52 @@ def _initialize_adapter_params(
Returns:
dict: Updated model variables with initialized adapter parameters.
"""
for i in range(num_layers):
layer_key = f"x_layers_{i}"
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"ff_layer"][ff_layer_key]["linear"]
original_w = linear["w"]
input_dim, output_dim = original_w.shape
std_dev = 1 / jnp.sqrt(lora_rank)
for i in range(num_layers):
layer_key = f"x_layers_{i}"
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"ff_layer"
][ff_layer_key]["linear"]
original_w = linear["w"]
input_dim, output_dim = original_w.shape
std_dev = 1 / jnp.sqrt(lora_rank)
normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(jax.random.key(seed),
(input_dim, lora_rank), jnp.float32)
lora_b = jnp.zeros((output_dim, lora_rank))
normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(
jax.random.key(seed), (input_dim, lora_rank), jnp.float32
)
lora_b = jnp.zeros((output_dim, lora_rank))
linear["lora_a"] = lora_a
linear["lora_b"] = lora_b
linear["lora_a"] = lora_a
linear["lora_b"] = lora_b
if use_dora:
norm = jnp.linalg.norm(original_w, ord=2, axis=0, keepdims=True)
linear["dora_m"] = norm
if use_dora:
norm = jnp.linalg.norm(original_w, ord=2, axis=0, keepdims=True)
linear["dora_m"] = norm
if lora_target_modules in ["all", "attention"]:
attention = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"self_attention"]
if lora_target_modules in ["all", "attention"]:
attention = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"self_attention"
]
for component in ["key", "query", "value", "post"]:
original_w = attention[component]["w"]
w_dim = original_w.shape[0]
std_dev = 1 / jnp.sqrt(lora_rank)
for component in ["key", "query", "value", "post"]:
original_w = attention[component]["w"]
w_dim = original_w.shape[0]
std_dev = 1 / jnp.sqrt(lora_rank)
normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(jax.random.key(seed), (w_dim, lora_rank),
jnp.float32)
lora_b = jnp.zeros((w_dim, lora_rank))
normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(
jax.random.key(seed), (w_dim, lora_rank), jnp.float32
)
lora_b = jnp.zeros((w_dim, lora_rank))
attention[component]["lora_a"] = lora_a
attention[component]["lora_b"] = lora_b
attention[component]["lora_a"] = lora_a
attention[component]["lora_b"] = lora_b
if use_dora:
norm = jnp.linalg.norm(original_w, ord=2, axis=0,
keepdims=True).astype(jnp.float32)
attention[component]["dora_m"] = norm
return mdl_vars
if use_dora:
norm = jnp.linalg.norm(
original_w, ord=2, axis=0, keepdims=True
).astype(jnp.float32)
attention[component]["dora_m"] = norm
return mdl_vars
+5 -6
View File
@@ -43,11 +43,11 @@ flags.DEFINE_list(
)
flags.DEFINE_string(
"local_model_path", None,
"local_model_path",
None,
"Path to a local .safetensors model file. If provided, overrides Hugging Face download."
)
class TimeSeriesDataset(Dataset):
"""Dataset for time series data compatible with TimesFM."""
@@ -148,7 +148,7 @@ def get_model(load_weights: bool = False):
use_positional_embedding=False,
context_len=192,
)
if load_weights:
if FLAGS.local_model_path:
tfm_config = TimesFMConfig()
@@ -157,12 +157,11 @@ def get_model(load_weights: bool = False):
else:
repo_id = "google/timesfm-2.0-500m-pytorch"
tfm = TimesFm(hparams=hparams,
checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))
checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))
tfm_config = tfm._model_config
model = PatchedTimeSeriesDecoder(tfm_config)
checkpoint_path = path.join(snapshot_download(repo_id),
"torch_model.ckpt")
checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt")
loaded_checkpoint = torch.load(checkpoint_path, weights_only=True)
model.load_state_dict(loaded_checkpoint)
+5 -7
View File
@@ -25,13 +25,11 @@ from timesfm.timesfm_base import (
import sys
try:
from timesfm.timesfm_jax import TimesFmJax as TimesFm
from timesfm import data_loader
from timesfm.timesfm_jax import TimesFmJax as TimesFm
from timesfm import data_loader
print(f"Loaded Jax TimesFM, likely because python version is {sys.version}.")
print(f"Loaded Jax TimesFM, likely because python version is {sys.version}.")
except Exception as _:
from timesfm.timesfm_torch import TimesFmTorch as TimesFm
from timesfm.timesfm_torch import TimesFmTorch as TimesFm
print(
f"Loaded PyTorch TimesFM, likely because python version is {sys.version}."
)
print(f"Loaded PyTorch TimesFM, likely because python version is {sys.version}.")
+14 -14
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Directory to extract time covariates.
Extract time covariates from datetime.
@@ -35,6 +36,7 @@ from pandas.tseries.offsets import Easter
from sklearn.preprocessing import StandardScaler
from tqdm import tqdm
# This is 183 to cover half a year (in both directions), also for leap years
# + 17 as Eastern can be between March, 22 - April, 25
MAX_WINDOW = 183 + 17
@@ -48,8 +50,9 @@ def _distance_to_holiday(holiday):
index - pd.Timedelta(days=MAX_WINDOW),
index + pd.Timedelta(days=MAX_WINDOW),
)
assert (len(holiday_date) != 0 # pylint: disable=g-explicit-length-test
), f"No closest holiday for the date index {index} found."
assert (
len(holiday_date) != 0 # pylint: disable=g-explicit-length-test
), f"No closest holiday for the date index {index} found."
# It sometimes returns two dates if it is exactly half a year after the
# holiday. In this case, the smaller distance (182 days) is returned.
return (index - holiday_date[0]).days
@@ -57,19 +60,16 @@ def _distance_to_holiday(holiday):
return _distance_to_day
EasterSunday = Holiday("Easter Sunday",
month=1,
day=1,
offset=[Easter(), Day(0)])
EasterSunday = Holiday(
"Easter Sunday", month=1, day=1, offset=[Easter(), Day(0)]
)
NewYearsDay = Holiday("New Years Day", month=1, day=1)
SuperBowl = Holiday("Superbowl",
month=2,
day=1,
offset=DateOffset(weekday=SU(1)))
MothersDay = Holiday("Mothers Day",
month=5,
day=1,
offset=DateOffset(weekday=SU(2)))
SuperBowl = Holiday(
"Superbowl", month=2, day=1, offset=DateOffset(weekday=SU(1))
)
MothersDay = Holiday(
"Mothers Day", month=5, day=1, offset=DateOffset(weekday=SU(2))
)
IndependenceDay = Holiday("Independence Day", month=7, day=4)
ChristmasEve = Holiday("Christmas", month=12, day=24)
ChristmasDay = Holiday("Christmas", month=12, day=25)
+17 -10
View File
@@ -25,12 +25,12 @@ import pandas as pd
from utilsforecast.processing import make_future_dataframe
if TYPE_CHECKING:
from . import xreg_lib
Category = xreg_lib.Category
XRegMode = xreg_lib.XRegMode
from . import xreg_lib
Category = xreg_lib.Category
XRegMode = xreg_lib.XRegMode
else:
Category = int | str
XRegMode = str
Category = int | str
XRegMode = str
_TOL = 1e-6
DEFAULT_QUANTILES = (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9)
@@ -45,7 +45,7 @@ def moving_average(arr, window_size):
"""Calculates the moving average using NumPy's convolution function."""
# Pad with zeros to handle initial window positions
arr_padded = np.pad(arr, (window_size - 1, 0), "constant")
smoothed_arr = (np.convolve(arr_padded, np.ones(window_size), "valid") /
smoothed_arr = (np.convolve(arr_padded, np.ones(window_size), "valid") /
window_size)
return [smoothed_arr, arr - smoothed_arr]
@@ -57,11 +57,18 @@ def freq_map(freq: str):
return 1
elif freq.endswith(("H", "T", "MIN", "D", "B", "U", "S")):
return 0
elif (freq.endswith(("W", "M")) or freq.startswith("W-") or
(freq.startswith("M") and len(freq) == 2)):
elif (
freq.endswith(("W", "M"))
or freq.startswith("W-")
or (freq.startswith("M") and len(freq) == 2)
):
return 1
elif (freq.endswith(("Y", "Q", "A")) or freq.startswith("Y-") or
freq.startswith("Q-") or freq.startswith("A-")):
elif (
freq.endswith(("Y", "Q", "A"))
or freq.startswith("Y-")
or freq.startswith("Q-")
or freq.startswith("A-")
):
return 2
else:
raise ValueError(f"Invalid frequency: {freq}")
+43 -41
View File
@@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from datetime import datetime, timedelta
import numpy as np
@@ -21,10 +22,10 @@ import pytest
import timesfm
def create_sample_dataframe(start_date: datetime,
end_date: datetime,
freq: str = "D") -> pd.DataFrame:
"""
def create_sample_dataframe(
start_date: datetime, end_date: datetime, freq: str = "D"
) -> pd.DataFrame:
"""
Create a sample DataFrame with time series data.
Args:
@@ -35,10 +36,10 @@ def create_sample_dataframe(start_date: datetime,
Returns:
pd.DataFrame: DataFrame with columns 'unique_id', 'ds', and 'ts'.
"""
date_range = pd.date_range(start=start_date, end=end_date, freq=freq)
ts_data = np.random.randn(len(date_range))
df = pd.DataFrame({"unique_id": "ts-1", "ds": date_range, "ts": ts_data})
return df
date_range = pd.date_range(start=start_date, end=end_date, freq=freq)
ts_data = np.random.randn(len(date_range))
df = pd.DataFrame({"unique_id": "ts-1", "ds": date_range, "ts": ts_data})
return df
@pytest.mark.parametrize("context_length", [128, 256, 512])
@@ -49,41 +50,42 @@ def test_timesfm_forecast_on_df(
prediction_length: int,
freq: str,
) -> None:
model = timesfm.TimesFm(
context_len=context_length,
horizon_len=prediction_length,
input_patch_len=32,
output_patch_len=128,
num_layers=20,
model_dims=1280,
backend="cpu",
)
model.load_from_checkpoint(repo_id="google/timesfm-1.0-200m")
model = timesfm.TimesFm(
context_len=context_length,
horizon_len=prediction_length,
input_patch_len=32,
output_patch_len=128,
num_layers=20,
model_dims=1280,
backend="cpu",
)
model.load_from_checkpoint(repo_id="google/timesfm-1.0-200m")
end_date = datetime.now()
start_date = end_date - timedelta(days=context_length)
input_df = create_sample_dataframe(start_date, end_date, freq)
end_date = datetime.now()
start_date = end_date - timedelta(days=context_length)
input_df = create_sample_dataframe(start_date, end_date, freq)
forecast_df = model.forecast_on_df(
inputs=input_df,
freq=freq,
value_name="ts",
num_jobs=-1,
)
forecast_df = model.forecast_on_df(
inputs=input_df,
freq=freq,
value_name="ts",
num_jobs=-1,
)
assert (
len(forecast_df) == prediction_length
), f"Expected forecast length of {prediction_length}, but got {len(forecast_df)}"
assert ("timesfm" in forecast_df.columns
), "Forecast DataFrame should contain 'timesfm' column"
assert (
len(forecast_df) == prediction_length
), f"Expected forecast length of {prediction_length}, but got {len(forecast_df)}"
assert (
"timesfm" in forecast_df.columns
), "Forecast DataFrame should contain 'timesfm' column"
last_input_date = input_df["ds"].max()
first_forecast_date = forecast_df["ds"].min()
expected_first_forecast_date = last_input_date + pd.Timedelta(1, unit=freq)
assert (
first_forecast_date == expected_first_forecast_date
), f"Forecast should start from {expected_first_forecast_date}, but starts from {first_forecast_date}"
last_input_date = input_df["ds"].max()
first_forecast_date = forecast_df["ds"].min()
expected_first_forecast_date = last_input_date + pd.Timedelta(1, unit=freq)
assert (
first_forecast_date == expected_first_forecast_date
), f"Forecast should start from {expected_first_forecast_date}, but starts from {first_forecast_date}"
print(
f"Successful forecast with context_length={context_length}, prediction_length={prediction_length}, freq={freq}"
)
print(
f"Successful forecast with context_length={context_length}, prediction_length={prediction_length}, freq={freq}"
)