Skip to content

Commit 2630a5a

Browse files
committed
Add LiteRT-LM export support for torch backend
Adds to CausalLM, enabling export to the LiteRT-LM bundle format with prefill/decode signatures. Key components: - : PyTorch adapter wrapping CausalLM for litert_torch.signature() export. Handles KV cache stacking/unstacking and traceable forward_prefill/forward_decode. - : Full export pipeline using litert_torch.signature('prefill', ...).signature('decode', ...). Builds LlmMetadata protobuf with start/stop tokens and model type. - : Integration tests verifying end-to-end export + numerical correctness. - : Unit tests mocking litert_torch and builder dependencies. - : Torch backend LiteRT test support. Requirements: - NAME litert-torch SYNOPSIS litert-torch COMMAND COMMANDS COMMAND is one of the following: export_hf Exports HuggingFace Transformers model to tflite. added to requirements.txt. - PyTorch backend required (Keras >= 3.15). - Prefill returns only KV caches (no logits), per LiteRT-LM spec. Tested with Gemma tiny models on torch backend.
1 parent f83a7fe commit 2630a5a

7 files changed

Lines changed: 1039 additions & 0 deletions

File tree

keras_hub/src/models/causal_lm.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -425,6 +425,39 @@ def export_to_transformers(self, path):
425425

426426
export_to_safetensors(self, path)
427427

428+
def export_to_litertlm(
429+
self,
430+
path,
431+
backend_constraint=None,
432+
prefill_seq_len=None,
433+
**kwargs,
434+
):
435+
"""Export the full CausalLM model to LiteRT-LM format.
436+
437+
This exports the model with ``prefill`` and ``decode`` signatures
438+
required by the LiteRT-LM executor, bundles the SentencePiece
439+
tokenizer, and writes an ``LlmMetadata`` protobuf into the
440+
`.litertlm` artifact.
441+
442+
Args:
443+
path: str. Path to save the `.litertlm` file.
444+
backend_constraint: Optional LiteRT-LM backend constraint, such as
445+
`"cpu"` or `"gpu"`.
446+
prefill_seq_len: int. Sequence length used when tracing the prefill
447+
signature. Defaults to the model's maximum sequence length.
448+
**kwargs: Additional kwargs forwarded to ``litert_torch``
449+
conversion.
450+
"""
451+
from keras_hub.src.utils.litertlm.export import export_to_litertlm
452+
453+
return export_to_litertlm(
454+
self,
455+
path,
456+
backend_constraint=backend_constraint,
457+
prefill_seq_len=prefill_seq_len,
458+
**kwargs,
459+
)
460+
428461
def _post_quantize(self, mode, **kwargs):
429462
super()._post_quantize(mode, **kwargs)
430463
# Reset the compiled generate function.

keras_hub/src/models/task_test.py

Lines changed: 180 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import os
22
import pathlib
3+
import types
34

45
import keras
56
import numpy as np
@@ -21,6 +22,7 @@
2122
from keras_hub.src.models.text_classifier import TextClassifier
2223
from keras_hub.src.tests.test_case import TestCase
2324
from keras_hub.src.tokenizers.tokenizer import Tokenizer
25+
from keras_hub.src.utils.litertlm import export as litertlm_export
2426
from keras_hub.src.utils.preset_utils import CONFIG_FILE
2527
from keras_hub.src.utils.preset_utils import METADATA_FILE
2628
from keras_hub.src.utils.preset_utils import MODEL_WEIGHTS_FILE
@@ -354,3 +356,181 @@ def test_export_missing_tokenizer(self):
354356
)
355357
with self.assertRaises(ValueError):
356358
causal_lm.export_to_transformers(export_path)
359+
360+
def test_export_to_litertlm(self):
361+
causal_lm, _ = self._create_gemma_for_export_tests()
362+
export_path = os.path.join(self.get_temp_dir(), "model.litertlm")
363+
364+
class FakeBuilder:
365+
instances = []
366+
367+
def __init__(self):
368+
self.tflite_model_args = None
369+
self.tokenizer_path = None
370+
self.tokenizer_path_exists = False
371+
self.llm_metadata_path = None
372+
type(self).instances.append(self)
373+
374+
def add_system_metadata(self, metadata):
375+
return self
376+
377+
def add_tflite_model(
378+
self,
379+
tflite_model_path,
380+
model_type,
381+
backend_constraint=None,
382+
):
383+
self.tflite_model_args = (
384+
tflite_model_path,
385+
model_type,
386+
backend_constraint,
387+
)
388+
return self
389+
390+
def add_sentencepiece_tokenizer(self, sp_tokenizer_path):
391+
self.tokenizer_path = sp_tokenizer_path
392+
self.tokenizer_path_exists = os.path.exists(sp_tokenizer_path)
393+
return self
394+
395+
def add_llm_metadata(self, llm_metadata_path):
396+
self.llm_metadata_path = llm_metadata_path
397+
return self
398+
399+
def build(self, stream):
400+
stream.write(b"litertlm")
401+
402+
fake_builder_module = types.SimpleNamespace(
403+
LitertLmFileBuilder=FakeBuilder,
404+
TfLiteModelType=types.SimpleNamespace(PREFILL_DECODE="prefill"),
405+
Metadata=types.SimpleNamespace,
406+
DType=types.SimpleNamespace(STRING="string"),
407+
)
408+
409+
class FakeEdgeModel:
410+
def export(self, path):
411+
with open(path, "wb") as f:
412+
f.write(b"tflite")
413+
414+
class FakeConverter:
415+
def signature(self, name, module, sample_kwargs=None, **kwargs):
416+
return self
417+
418+
def convert(self):
419+
return FakeEdgeModel()
420+
421+
fake_litert_torch = types.SimpleNamespace(
422+
signature=FakeConverter().signature,
423+
)
424+
425+
with pytest.MonkeyPatch.context() as mp:
426+
mp.setattr(keras.config, "backend", lambda: "torch")
427+
mp.setattr(
428+
litertlm_export,
429+
"_import_litert_lm_builder",
430+
lambda: fake_builder_module,
431+
)
432+
mp.setattr(
433+
litertlm_export,
434+
"litert_torch",
435+
fake_litert_torch,
436+
)
437+
causal_lm.export_to_litertlm(export_path, backend_constraint="cpu")
438+
439+
self.assertTrue(os.path.exists(export_path))
440+
self.assertEqual(len(FakeBuilder.instances), 1)
441+
builder = FakeBuilder.instances[0]
442+
self.assertEqual(builder.tflite_model_args[2], "cpu")
443+
self.assertTrue(builder.tokenizer_path.endswith("vocabulary.spm"))
444+
self.assertTrue(builder.tokenizer_path_exists)
445+
self.assertIsNotNone(builder.llm_metadata_path)
446+
447+
def test_export_to_litertlm_after_keras_save_load(self):
448+
causal_lm, _ = self._create_gemma_for_export_tests()
449+
keras_path = os.path.join(self.get_temp_dir(), "model.keras")
450+
export_path = os.path.join(self.get_temp_dir(), "model.litertlm")
451+
causal_lm.save(keras_path)
452+
restored = keras.saving.load_model(keras_path)
453+
454+
class FakeBuilder:
455+
instances = []
456+
457+
def __init__(self):
458+
self.tokenizer_path_exists = False
459+
type(self).instances.append(self)
460+
461+
def add_system_metadata(self, metadata):
462+
return self
463+
464+
def add_tflite_model(
465+
self,
466+
tflite_model_path,
467+
model_type,
468+
backend_constraint=None,
469+
):
470+
return self
471+
472+
def add_sentencepiece_tokenizer(self, sp_tokenizer_path):
473+
self.tokenizer_path = sp_tokenizer_path
474+
self.tokenizer_path_exists = os.path.exists(sp_tokenizer_path)
475+
return self
476+
477+
def add_llm_metadata(self, llm_metadata_path):
478+
return self
479+
480+
def build(self, stream):
481+
stream.write(b"litertlm")
482+
483+
fake_builder_module = types.SimpleNamespace(
484+
LitertLmFileBuilder=FakeBuilder,
485+
TfLiteModelType=types.SimpleNamespace(PREFILL_DECODE="prefill"),
486+
Metadata=types.SimpleNamespace,
487+
DType=types.SimpleNamespace(STRING="string"),
488+
)
489+
490+
class FakeEdgeModel:
491+
def export(self, path):
492+
with open(path, "wb") as f:
493+
f.write(b"tflite")
494+
495+
class FakeConverter:
496+
def signature(self, name, module, sample_kwargs=None, **kwargs):
497+
return self
498+
499+
def convert(self):
500+
return FakeEdgeModel()
501+
502+
fake_litert_torch = types.SimpleNamespace(
503+
signature=FakeConverter().signature,
504+
)
505+
506+
with pytest.MonkeyPatch.context() as mp:
507+
mp.setattr(keras.config, "backend", lambda: "torch")
508+
mp.setattr(
509+
litertlm_export,
510+
"_import_litert_lm_builder",
511+
lambda: fake_builder_module,
512+
)
513+
mp.setattr(
514+
litertlm_export,
515+
"litert_torch",
516+
fake_litert_torch,
517+
)
518+
restored.export_to_litertlm(export_path)
519+
520+
self.assertTrue(os.path.exists(export_path))
521+
self.assertEqual(len(FakeBuilder.instances), 1)
522+
self.assertTrue(FakeBuilder.instances[0].tokenizer_path_exists)
523+
524+
def test_export_to_litertlm_rejects_non_sentencepiece_tokenizer(self):
525+
causal_lm, preprocessor = self._create_gemma_for_export_tests()
526+
export_path = os.path.join(self.get_temp_dir(), "model.litertlm")
527+
528+
class UnsupportedTokenizer(Tokenizer):
529+
def __init__(self):
530+
super().__init__()
531+
self.file_assets = ["vocabulary.json"]
532+
533+
preprocessor.tokenizer = UnsupportedTokenizer()
534+
535+
with self.assertRaises(ValueError):
536+
causal_lm.export_to_litertlm(export_path)
Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
from keras_hub.src.utils.litertlm.export import export_to_litertlm
2+
3+
__all__ = ["export_to_litertlm"]

0 commit comments

Comments
 (0)