|
1 | 1 | import os |
2 | 2 | import pathlib |
| 3 | +import types |
3 | 4 |
|
4 | 5 | import keras |
5 | 6 | import numpy as np |
|
21 | 22 | from keras_hub.src.models.text_classifier import TextClassifier |
22 | 23 | from keras_hub.src.tests.test_case import TestCase |
23 | 24 | from keras_hub.src.tokenizers.tokenizer import Tokenizer |
| 25 | +from keras_hub.src.utils.litertlm import export as litertlm_export |
24 | 26 | from keras_hub.src.utils.preset_utils import CONFIG_FILE |
25 | 27 | from keras_hub.src.utils.preset_utils import METADATA_FILE |
26 | 28 | from keras_hub.src.utils.preset_utils import MODEL_WEIGHTS_FILE |
@@ -354,3 +356,181 @@ def test_export_missing_tokenizer(self): |
354 | 356 | ) |
355 | 357 | with self.assertRaises(ValueError): |
356 | 358 | 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) |
0 commit comments