File size: 6,648 Bytes
60b21d3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 | # SPDX-FileCopyrightText: 2025 Stanford University, ETH Zurich, and the project authors (see CONTRIBUTORS.md)
# SPDX-FileCopyrightText: 2025 This source file is part of the OpenTSLM open-source project.
#
# SPDX-License-Identifier: MIT
import numpy as np
from functools import partial
from dataclasses import dataclass
# Please check the original code at https://github.com/ngruver/llmtime/blob/main/data/serialize.py
@dataclass
class SerializerSettings:
"""
Settings for time series serialization and deserialization.
"""
base: int = 10 # Numeric base for representation
prec: int = 3 # Number of digits after the 'decimal' point
signed: bool = True # Whether to include a sign for positive values
fixed_length: bool = False # If True, pad to a fixed number of digits
max_val: float = 1e7 # Maximum absolute value allowed
time_sep: str = ' ,' # Separator between time steps
bit_sep: str = ' ' # Separator between individual digits
plus_sign: str = '' # String to prepend for positive values
minus_sign: str = ' -' # String to prepend for negative values
half_bin_correction: bool = True # Apply half-bin correction on deserialization
decimal_point: str = '' # Literal to mark decimal point in serialized string
missing_str: str = ' Nan' # Representation for missing (NaN) values
def vec_num2repr(val: np.ndarray, base: int, prec: int, max_val: float):
"""
Convert an array of floats into sign and digit representation in the given base.
"""
base = float(base)
sign = np.where(val >= 0, 1, -1)
mag = np.abs(val)
max_bit_pos = int(np.ceil(np.log(max_val) / np.log(base)).item())
# Compute digits before the "decimal"
before = []
rem = mag.copy()
for i in range(max_bit_pos):
power = base ** (max_bit_pos - i - 1)
digit = (rem / power).astype(int)
before.append(digit)
rem = rem - digit * power
before = np.stack(before, axis=-1)
# Compute digits after the "decimal"
if prec > 0:
after = []
for i in range(prec):
power = base ** (-(i + 1))
digit = (rem / power).astype(int)
after.append(digit)
rem = rem - digit * power
after = np.stack(after, axis=-1)
digits = np.concatenate([before, after], axis=-1)
else:
digits = before
return sign, digits
def vec_repr2num(sign: np.ndarray, digits: np.ndarray, base: int, prec: int, half_bin_correction: bool = True) -> np.ndarray:
"""
Convert sign and digit arrays back into floats.
"""
base = float(base)
# Reverse the digit order for positional weights
flipped = np.flip(digits, axis=-1)
total_digits = digits.shape[1]
# Compute positional powers
powers = -np.arange(-prec, total_digits-prec)
mags = np.sum(flipped / (base ** powers), axis=-1)
if half_bin_correction:
mags = mags + (0.5 / (base ** prec))
return sign * mags
def serialize_arr(arr: np.ndarray, settings: SerializerSettings) -> str:
"""
Serialize a 1D numpy array into a digit-level string using the given settings.
"""
# Validate range
clean = np.where(np.isnan(arr), 0.0, arr)
assert np.all(np.abs(clean) <= settings.max_val), \
f"Values must be within ±{settings.max_val}"
# Convert numbers to sign and digits
to_repr = partial(vec_num2repr, base=settings.base, prec=settings.prec, max_val=settings.max_val)
sign_arr, digits_arr = to_repr(clean)
is_nan = np.isnan(arr)
def format_digits(digits):
# Optionally strip leading zeros
if not settings.fixed_length:
nz = np.where(digits != 0)[0]
if nz.size > 0:
digits = digits[nz[0]:]
else:
digits = np.array([0], dtype=int)
# Insert decimal point if specified
if settings.decimal_point and settings.prec > 0:
point_idx = len(digits) - settings.prec
digits = np.concatenate([digits[:point_idx], [-1], digits[point_idx:]])
# Join with bit separator
return settings.bit_sep.join(str(d) for d in digits if d != -1)
parts = []
for s, digs, missing in zip(sign_arr, digits_arr, is_nan):
if missing:
parts.append(settings.missing_str)
else:
sign_str = settings.plus_sign if s > 0 else settings.minus_sign
digit_str = format_digits(digs)
parts.append(sign_str + digit_str)
result = settings.time_sep.join(parts) + settings.time_sep
return result
def deserialize_str(bit_str: str, settings: SerializerSettings, ignore_last: bool = False, steps: int = None) -> np.ndarray:
"""
Deserialize a serialized string back into a numpy array of floats.
"""
tokens = [t for t in bit_str.split(settings.time_sep) if t]
if ignore_last:
tokens = tokens[:-1]
if steps is not None:
tokens = tokens[:steps]
sign_list = []
digit_list = []
for tok in tokens:
if tok == settings.missing_str.strip():
sign_list.append(1)
digit_list.append([0])
continue
# Determine sign
if settings.signed and tok.startswith(settings.minus_sign):
sign_list.append(-1)
tok = tok[len(settings.minus_sign):]
else:
sign_list.append(1)
if settings.signed and tok.startswith(settings.plus_sign):
tok = tok[len(settings.plus_sign):]
# Split digits
digs = [int(ch) for ch in tok.split(settings.bit_sep) if ch.isdigit()]
digit_list.append(digs)
# Pad to equal length
maxlen = max(len(d) for d in digit_list)
padded = np.array([([0]*(maxlen - len(d)) + d) for d in digit_list])
nums = vec_repr2num(
np.array(sign_list), padded,
base=settings.base, prec=settings.prec,
half_bin_correction=settings.half_bin_correction
)
return nums
gpt_settings = SerializerSettings(base=10, prec=3, signed=True, time_sep=', ', bit_sep=' ', minus_sign='-')
llama = SerializerSettings(base=10, prec=3, signed=True, half_bin_correction=True)
# Then wrap serialize_arr so it only takes the array
def gpt_formatter(arr: np.ndarray) -> str:
return serialize_arr(arr, gpt_settings)
def llama_formatter(arr: np.ndarray) -> str:
return serialize_arr(arr, llama)
# Backward compatibility: default to llama_formatter
# (or you can set to gpt_formatter if you prefer)
gruver_et_al_formatter = llama_formatter |