bertin-project/bertin-base-random-exp-512seqlen
025
1#!/usr/bin/env python2import tempfile3 4import jax5from jax import numpy as jnp6from transformers import AutoTokenizer, FlaxRobertaForMaskedLM, RobertaForMaskedLM7 8 9def to_f32(t):10 return jax.tree_map(lambda x: x.astype(jnp.float32) if x.dtype == jnp.bfloat16 else x, t)11 12 13def main():14 # Saving extra files from config.json and tokenizer.json files15 tokenizer = AutoTokenizer.from_pretrained("./")16 tokenizer.save_pretrained("./")17 18 # Temporary saving bfloat16 Flax model into float3219 tmp = tempfile.mkdtemp()20 flax_model = FlaxRobertaForMaskedLM.from_pretrained("./")21 flax_model.params = to_f32(flax_model.params)22 flax_model.save_pretrained(tmp)23 # Converting float32 Flax to PyTorch24 model = RobertaForMaskedLM.from_pretrained(tmp, from_flax=True)25 model.save_pretrained("./", save_config=False)26 27 28if __name__ == "__main__":29 main()30 