Spaces:
Sleeping
Sleeping
| # model_build.py | |
| import os | |
| os.environ["KERAS_BACKEND"] = "jax" | |
| import keras | |
| from veylon_model import create_llm | |
| from tokenizer import TokenizerWrapper | |
| from config import ( | |
| CONTEXT, | |
| ffn_mult, | |
| d_Latent, | |
| D_MODEL, | |
| numberofheads, | |
| numberoflayers, | |
| vocab_size | |
| ) | |
| # Mixed precision is set inside veylon_model.py, but we ensure it here too | |
| keras.mixed_precision.set_global_policy("mixed_bfloat16") | |
| tokenizer = TokenizerWrapper("tokenizer.json") | |
| # n_kv_heads is completely removed - pure MLA | |
| model = create_llm( | |
| vocab_size=vocab_size, | |
| d_model=D_MODEL, | |
| n_layers=numberoflayers, | |
| n_heads=numberofheads, | |
| d_latent=d_Latent, | |
| ffn_mult=ffn_mult, | |
| max_seq_len=CONTEXT, | |
| use_moe=False, | |
| ) | |
| optimizer = keras.optimizers.AdamW( | |
| learning_rate=1e-4, | |
| weight_decay=0.01, | |
| global_clipnorm=1.0, | |
| ) | |
| loss_fn = keras.losses.SparseCategoricalCrossentropy( | |
| from_logits=True, | |
| ) | |
| model.compile( | |
| optimizer=optimizer, | |
| loss=loss_fn, | |
| jit_compile=True, | |
| ) | |
| # Explicitly build to print summary | |
| model.build((None, CONTEXT)) | |
| model.summary() |