EasyDeL is an open-source framework designed to enhance and streamline the training process of machine learning
models. With a primary focus on Jax, EasyDeL aims to provide convenient and effective solutions for
training Flax/Jax models on TPU/GPU, for both serving and training purposes.
1from easydel import AutoEasyDeLModelForCausalLM
2from jax import numpy as jnp, lax
3
4model = AutoEasyDeLModelForCausalLM.from_pretrained(
5 f"REPO_ID/BaseTrainer",
6 dtype=...,
7 param_dtype=...,
8 precision=lax.Precision("fastest"),
9 auto_shard_model=True,
10)
1# Partition Rules
2( ('model/embed_tokens/embedding', PartitionSpec('tp', ('fsdp', 'sp'))),
3 ( 'self_attn/(q_proj|k_proj|v_proj)/kernel',
4 PartitionSpec(('fsdp', 'sp'), 'tp')),
5 ('self_attn/o_proj/kernel', PartitionSpec('tp', ('fsdp', 'sp'))),
6 ('mlp/gate_proj/kernel', PartitionSpec(('fsdp', 'sp'), 'tp')),
7 ('mlp/down_proj/kernel', PartitionSpec('tp', ('fsdp', 'sp'))),
8 ('mlp/up_proj/kernel', PartitionSpec(('fsdp', 'sp'), 'tp')),
9 ('input_layernorm/kernel', PartitionSpec(None,)),
10 ('post_attention_layernorm/kernel', PartitionSpec(None,)),
11 ('model/norm/kernel', PartitionSpec(None,)),
12 ('lm_head/kernel', PartitionSpec(('fsdp', 'sp'), 'tp')),
13 ('.*', PartitionSpec(None,)))