Download tests/test_model.py from hermescures1/splitbit-llm: direct link, hf CLI and curl.
- Browser
- Download file 2.99 kB
-
https://huggingface.co/hermescures1/splitbit-llm/resolve/main/tests/test_model.py
- Command line
-
hf download hf://hermescures1/splitbit-llm/tests/test_model.py
-
curl -L -o test_model.py https://huggingface.co/hermescures1/splitbit-llm/resolve/main/tests/test_model.py
2.99 kB
| """Test model: forward pass, generation, save/load.""" | |
| import sys | |
| import os | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) | |
| import numpy as np | |
| from splitbit_llm.config import get_model_config, HardwareTier | |
| from splitbit_llm.model.model import SplitBitLLM | |
| from splitbit_llm.model.tokenizer import BPETokenizer | |
| def test_model_forward(): | |
| """Test model forward pass.""" | |
| cfg = get_model_config(HardwareTier.MOBILE) | |
| cfg.vocab_size = 256 | |
| model = SplitBitLLM(config=cfg) | |
| token_ids = np.array([[1, 5, 10, 15, 20]], dtype=np.int64) | |
| logits = model.forward(token_ids) | |
| assert logits.shape == (1, 5, 256), f"Wrong shape: {logits.shape}" | |
| print(f" Logits shape: {logits.shape}") | |
| print(f" Param count: {model.param_count:,}") | |
| def test_model_generate(): | |
| """Test text generation.""" | |
| cfg = get_model_config(HardwareTier.MOBILE) | |
| cfg.vocab_size = 256 | |
| model = SplitBitLLM(config=cfg) | |
| output = model.generate("Hello", max_tokens=10, temperature=0.7) | |
| assert isinstance(output, str), f"Expected str, got {type(output)}" | |
| assert len(output) > 0, "Empty output" | |
| print(f" Generated: {repr(output[:50])}") | |
| def test_model_generate_stream(): | |
| """Test streaming generation.""" | |
| cfg = get_model_config(HardwareTier.MOBILE) | |
| cfg.vocab_size = 256 | |
| model = SplitBitLLM(config=cfg) | |
| chunks = list(model.generate_stream("Hello", max_tokens=10, temperature=0.7)) | |
| assert len(chunks) > 0, "No chunks generated" | |
| print(f" Chunks: {len(chunks)}") | |
| def test_model_with_tokenizer(): | |
| """Test model with trained tokenizer.""" | |
| tok = BPETokenizer(vocab_size=256) | |
| tok.train("Hello world! This is a test. Hello world again. The quick brown fox jumps.") | |
| cfg = get_model_config(HardwareTier.MOBILE) | |
| cfg.vocab_size = 256 | |
| model = SplitBitLLM(config=cfg, tokenizer=tok) | |
| output = model.generate("Hello", max_tokens=10, temperature=0.7) | |
| assert isinstance(output, str) | |
| print(f" Generated with tokenizer: {repr(output[:50])}") | |
| def test_model_truncation(): | |
| """Test that long inputs are truncated to max_seq_len.""" | |
| cfg = get_model_config(HardwareTier.MOBILE) | |
| cfg.vocab_size = 256 | |
| cfg.max_seq_len = 32 | |
| model = SplitBitLLM(config=cfg) | |
| # Input longer than max_seq_len | |
| long_input = np.array([[i for i in range(100)]], dtype=np.int64) | |
| logits = model.forward(long_input) | |
| assert logits.shape[1] == 32, f"Should truncate to 32, got {logits.shape[1]}" | |
| print(f" Truncated to: {logits.shape[1]}") | |
| if __name__ == "__main__": | |
| print("Running model tests...") | |
| test_model_forward() | |
| print(" β test_model_forward") | |
| test_model_generate() | |
| print(" β test_model_generate") | |
| test_model_generate_stream() | |
| print(" β test_model_generate_stream") | |
| test_model_with_tokenizer() | |
| print(" β test_model_with_tokenizer") | |
| test_model_truncation() | |
| print(" β test_model_truncation") | |
| print("\nAll model tests passed!") | |