deepmodeling / deepmodeling/DMFF

[Feature Request] Refactor the cpp interface of the saved DMFF jax model with MD engine

Open
#173 0 comments 0 reactions 0 assignees View on GitHub
enhancement
Dominant language
Python
Stars
198
Forks
49
PR merge metrics
No merged PRs in 30d

Description

### Summary

Moving the `jax2tf` to `HLOModule` for the cpp interface of the saved DMFF model.

### Motivation

The current implementation of the cpp interface between the saved DMFF model and MD engine was based on the [jax2tf](https://github.com/google/jax/tree/jaxlib-v0.4.25/jax/experimental/jax2tf).
The `jax2tf` was used to convert the the jax function to TensorFlow function.
However, as an experimental feature of JAX, `jax2tf` does have some limitations for production use.

1. Limited support for custom calls. https://github.com/google/jax/tree/jaxlib-v0.4.25/jax/experimental/jax2tf#native-serialization-supports-only-select-custom-calls.
Occurred when using JAX 0.4.24 + TF 2.15/2.14
2. Unsupported data type f64, s64,

![image](https://github.com/deepmodeling/DMFF/assets/13417572/5fdf94d8-781c-4f38-95d3-51d6249be6b1)

### Suggested Solutions

https://github.com/google/jax/issues/1871
Old solution.

Lack of documentation, more exploration required

### Further Information, Files, and Links

_No response_

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.