diff --git a/convert.py b/convert.py index bfc600f..e34a4d0 100644 --- a/convert.py +++ b/convert.py @@ -8,7 +8,6 @@ (".q.", ".query_proj."), (".v.", ".value_proj."), ("shared.", "wte."), - ("lm_head.", "lm_head.linear."), (".layer.0.layer_norm.", ".ln1."), (".layer.1.layer_norm.", ".ln2."), (".layer.2.layer_norm.", ".ln3."), diff --git a/mlx_bitnet.py b/mlx_bitnet.py index 6eea92f..b02c482 100644 --- a/mlx_bitnet.py +++ b/mlx_bitnet.py @@ -846,6 +846,7 @@ def sanitize_config(_config: BitnetConfig) -> MinimalBitnetConfig: intermediate_size=_config.intermediate_size, max_position_embeddings=_config.max_position_embeddings, num_attention_heads=_config.num_attention_heads, + num_hidden_layers=_config.num_hidden_layers, num_key_value_heads=_config.num_key_value_heads, pad_token_id=_config.pad_token_id, rms_norm_eps=_config.rms_norm_eps, diff --git a/run_mlx.py b/run_mlx.py new file mode 100644 index 0000000..93543aa --- /dev/null +++ b/run_mlx.py @@ -0,0 +1,31 @@ +import argparse +import mlx.core as mx +from mlx_bitnet import load_causal_model + +parser = argparse.ArgumentParser() +parser.add_argument("--model", default="1bitLLM/bitnet_b1_58-xl", type=str) +parser.add_argument("--prompt", default="The meaning of life is", type=str) +parser.add_argument("--max_tokens", default=50, type=int) +parser.add_argument("--temp", default=0.0, type=float) +args = parser.parse_args() + +print(f"Loading model {args.model} ...") +model, tokenizer = load_causal_model(args.model) + +tokens = tokenizer.encode(args.prompt) +input_ids = mx.array([tokens]) +attention_mask = mx.ones_like(input_ids) + +generated = list(tokens) +decoded_so_far = tokenizer.decode(tokens) +print(args.prompt, end="", flush=True) + +for token in model.generate(input_ids, attention_mask, temp=args.temp): + generated.append(token.item()) + new_text = tokenizer.decode(generated) + print(new_text[len(decoded_so_far):], end="", flush=True) + decoded_so_far = new_text + if len(generated) - len(tokens) >= args.max_tokens: + break + +print() diff --git a/test_interop.py b/test_interop.py index ee0e1bc..67de7dd 100644 --- a/test_interop.py +++ b/test_interop.py @@ -26,7 +26,6 @@ from torch_bitnet import BitnetForCausalLM as TorchBitnetForCausalLM from torch_bitnet import BitnetDecoderLayer as TorchBitnetDecoderLayer from transformers.activations import silu as torch_silu -from training.bit_linear import weight_quant as bit_linear_weight_quant class TestBitLinearInterop(unittest.TestCase): def setUp(self): diff --git a/torch_bitnet.py b/torch_bitnet.py index ef58099..a77456c 100644 --- a/torch_bitnet.py +++ b/torch_bitnet.py @@ -585,7 +585,7 @@ def _update_causal_mask(self, attention_mask, input_tensor, cache_position): class BitnetForCausalLM(BitnetPreTrainedModel): - _tied_weights_keys = ["lm_head.weight"] + _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} def __init__(self, config): super().__init__(config)