()
| 290 | |
| 291 | |
| 292 | def main(): |
| 293 | parser = ArgumentParser() |
| 294 | parser.add_argument( |
| 295 | "--use-sdpa-with-kv-cache", |
| 296 | default=True, |
| 297 | action=BooleanOptionalAction, |
| 298 | help="Use sdpa_with_kv_cache custom op in LLava text model.", |
| 299 | ) |
| 300 | parser.add_argument( |
| 301 | "--max-context-len", |
| 302 | required=True, |
| 303 | type=int, |
| 304 | help="Maximum context length for the text model.", |
| 305 | ) |
| 306 | parser.add_argument( |
| 307 | "--max-seq-len", |
| 308 | default=768, |
| 309 | type=int, |
| 310 | help="Maximum sequence length for the text model.", |
| 311 | ) |
| 312 | parser.add_argument( |
| 313 | "--pte-name", |
| 314 | default="llava_combined_xnnpack.pte", |
| 315 | help="Name of the exported ExecuTorch program.", |
| 316 | ) |
| 317 | parser.add_argument( |
| 318 | "--with-artifacts", |
| 319 | default=False, |
| 320 | action=BooleanOptionalAction, |
| 321 | help="Generate artifacts for llava runner.", |
| 322 | ) |
| 323 | parser.add_argument( |
| 324 | "--profile_memory", |
| 325 | required=False, |
| 326 | action="store_true", |
| 327 | help="Generate chrome trace of activation memory for intermediate tensors.", |
| 328 | ) |
| 329 | args = parser.parse_args() |
| 330 | |
| 331 | # Create LlmConfig from args |
| 332 | llm_config = create_llava_config_from_args(args) |
| 333 | |
| 334 | logging.info( |
| 335 | f"Exporting Llava model to ExecuTorch with sdpa_with_kv_cache: {llm_config.model.use_sdpa_with_kv_cache}, max_seq_len: {llm_config.export.max_seq_length}, max_context_len: {llm_config.export.max_context_length}" |
| 336 | ) |
| 337 | |
| 338 | llava_model = LlavaModel( |
| 339 | use_sdpa_with_kv_cache_op=llm_config.model.use_sdpa_with_kv_cache, |
| 340 | max_seq_len=llm_config.export.max_seq_length, |
| 341 | max_context_len=llm_config.export.max_context_length, |
| 342 | ) |
| 343 | |
| 344 | executorch_program = export_all(llava_model) |
| 345 | |
| 346 | # memory profiling |
| 347 | if llm_config.debug.profile_memory: |
| 348 | for method_name in executorch_program.methods: |
| 349 | generate_memory_trace( |
no test coverage detected