Wiola360M / configuration_wiola.py
aryuemaan's picture
Upload folder using huggingface_hub
2db32a1 verified
Raw
History Blame Contribute Delete
6.36 kB
# coding=utf-8
# Copyright 2025 The Wiola / OSCOWL-AI authors.
# Licensed under the Apache License, Version 2.0 (the "License").
"""Wiola model configuration."""
from transformers.configuration_utils import PretrainedConfig
class WiolaConfig(PretrainedConfig):
r"""
Configuration class for a :class:`WiolaForCausalLM` model.
This stores every hyper-parameter described in the Wiola paper. Defaults
correspond to the **wiola-360m** variant. The four published sizes are
available as YAML files under ``configs/`` and as named presets via
:meth:`WiolaConfig.from_preset`.
Args:
vocab_size (int): Vocabulary size of the BPE tokenizer.
hidden_size (int): Model hidden dimension ``d``.
num_hidden_layers (int): Number of decoder layers ``L``.
num_attention_heads (int): Number of query heads ``H``.
num_key_value_heads (int): Number of key/value heads ``H_kv`` (GQA).
max_position_embeddings (int): Maximum context length ``T``.
dsff_narrow_size (int): DSFF Stream A width ``d_A``.
dsff_wide_size (int): DSFF Stream B width ``d_B``.
srpe_theta (float): SRPE base theta ``theta_0``.
srpe_spiral_divisor (int): SRPE spiral divisor ``k_s``.
srpe_radial_amplitude (float): SRPE radial amplitude ``a_s``.
srpe_radial_frequency (float): SRPE radial frequency ``f_s``.
atm_threshold (float): ATM cosine-similarity merge threshold ``tau``.
atm_enabled (bool): Master switch for Adaptive Token Merging in training.
gcla_lookback (int): GCLA lookback depth ``Lambda``.
gcla_gate_init (float): Logit ``phi`` used to initialise the blend gate
``beta = sigmoid(phi)``.
rms_norm_eps (float): Epsilon for WiolaRMSNorm.
initializer_range (float): Stddev for truncated-normal init.
tie_word_embeddings (bool): Tie input embedding and LM head.
"""
model_type = "wiola"
keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
vocab_size: int = 32000,
hidden_size: int = 1024,
num_hidden_layers: int = 16,
num_attention_heads: int = 16,
num_key_value_heads: int = 4,
max_position_embeddings: int = 2048,
dsff_narrow_size: int = 1024,
dsff_wide_size: int = 4096,
srpe_theta: float = 10000.0,
srpe_spiral_divisor: int = 8,
srpe_radial_amplitude: float = 0.1,
srpe_radial_frequency: float = 0.01,
atm_threshold: float = 0.92,
atm_enabled: bool = True,
gcla_lookback: int = 2,
gcla_gate_init: float = -3.0,
rms_norm_eps: float = 1e-6,
initializer_range: float = 0.02,
use_cache: bool = True,
pad_token_id: int = 0,
bos_token_id: int = 1,
eos_token_id: int = 2,
tie_word_embeddings: bool = True,
**kwargs,
):
self.vocab_size = vocab_size
self.hidden_size = hidden_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
# GQA: default num_key_value_heads to num_attention_heads (MHA) when unset.
if num_key_value_heads is None:
num_key_value_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads
self.max_position_embeddings = max_position_embeddings
self.dsff_narrow_size = dsff_narrow_size
self.dsff_wide_size = dsff_wide_size
self.srpe_theta = srpe_theta
self.srpe_spiral_divisor = srpe_spiral_divisor
self.srpe_radial_amplitude = srpe_radial_amplitude
self.srpe_radial_frequency = srpe_radial_frequency
self.atm_threshold = atm_threshold
self.atm_enabled = atm_enabled
self.gcla_lookback = gcla_lookback
self.gcla_gate_init = gcla_gate_init
self.rms_norm_eps = rms_norm_eps
self.initializer_range = initializer_range
self.use_cache = use_cache
if hidden_size % num_attention_heads != 0:
raise ValueError(
f"hidden_size ({hidden_size}) must be divisible by "
f"num_attention_heads ({num_attention_heads})."
)
if num_attention_heads % num_key_value_heads != 0:
raise ValueError(
f"num_attention_heads ({num_attention_heads}) must be divisible "
f"by num_key_value_heads ({num_key_value_heads})."
)
super().__init__(
pad_token_id=pad_token_id,
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
tie_word_embeddings=tie_word_embeddings,
**kwargs,
)
@property
def head_dim(self) -> int:
return self.hidden_size // self.num_attention_heads
# Convenience presets matching the paper's model family. -----------------
_PRESETS = {
"wiola-120m": dict(
hidden_size=768,
num_hidden_layers=12,
num_attention_heads=12,
num_key_value_heads=4,
dsff_narrow_size=768,
dsff_wide_size=3072,
),
"wiola-360m": dict(
hidden_size=1024,
num_hidden_layers=16,
num_attention_heads=16,
num_key_value_heads=4,
dsff_narrow_size=1024,
dsff_wide_size=4096,
),
"wiola-700m": dict(
hidden_size=1536,
num_hidden_layers=24,
num_attention_heads=16,
num_key_value_heads=8,
dsff_narrow_size=1536,
dsff_wide_size=6144,
),
"wiola-1.5b": dict(
hidden_size=2048,
num_hidden_layers=28,
num_attention_heads=16,
num_key_value_heads=8,
dsff_narrow_size=2048,
dsff_wide_size=8192,
),
}
@classmethod
def from_preset(cls, name: str, **overrides) -> "WiolaConfig":
if name not in cls._PRESETS:
raise KeyError(f"Unknown preset '{name}'. Choose from {list(cls._PRESETS)}.")
params = dict(cls._PRESETS[name])
params.update(overrides)
return cls(**params)