distily/distily_profile_smollm
030
Summary
Distilled with Distily library using teacher model HuggingFaceTB/SmolLM-135M on dataset wikimedia/wikipedia.
<!-- This model card has been generated automatically according to the information the Trainer had access to. You should probably proofread and complete it, then remove this comment.
Model description
More information needed
Intended uses & limitations
More information needed -->
Model Architecture:
- Architecture:
LlamaForCausalLM - Total Parameters: 81,413,568
- Data Type (dtype): torch.bfloat16
- Model Size: 0.15 GB
<details> <summary>Student Model Details</summary>
LlamaForCausalLM(
(model): LlamaModel(
(embed_tokens): Embedding(49152, 576)
(layers): ModuleList(
(0-14): 15 x LlamaDecoderLayer(
(self_attn): LlamaSdpaAttention(
(q_proj): Linear(in_features=576, out_features=576, bias=False)
(k_proj): Linear(in_features=576, out_features=192, bias=False)
(v_proj): Linear(in_features=576, out_features=192, bias=False)
(o_proj): Linear(in_features=576, out_features=576, bias=False)
(rotary_emb): LlamaRotaryEmbedding()
)
(mlp): LlamaMLP(
(gate_proj): Linear(in_features=576, out_features=1536, bias=False)
(up_proj): Linear(in_features=576, out_features=1536, bias=False)
(down_proj): Linear(in_features=1536, out_features=576, bias=False)
(act_fn): SiLU()
)
(input_layernorm): LlamaRMSNorm((576,), eps=1e-05)
(post_attention_layernorm): LlamaRMSNorm((576,), eps=1e-05)
)
)
(norm): LlamaRMSNorm((576,), eps=1e-05)
(rotary_emb): LlamaRotaryEmbedding()
)
(lm_head): Linear(in_features=576, out_features=49152, bias=False)
)</details> <br/>
Resource Usage
- Max Train VRAM Use: 12.7946 GB
- Available VRAM: 23.4329 GB
- GPUs:
- 1x NVIDIA GeForce RTX 4090
- CPUs: 64
- CPU Memory: 251.7299 GB
- CPU Memory Bandwidth: 1600 GB/s
Distillation (Teacher -> Student) Architecture Difference:
- Architecture:
LlamaForCausalLM->LlamaForCausalLM - Total Parameters: 134,515,008 -> 81,413,568
- Data Type (dtype): torch.bfloat16 -> torch.bfloat16
- Model Size: 0.25 GB -> 0.15 GB
<details> <summary>Module Diff Details</summary>
--- teacher model modules
+++ student model modules
@@ -2,7 +2,7 @@
(model): LlamaModel(
(embed_tokens): Embedding(49152, 576)
(layers): ModuleList(
- (0-29): 30 x LlamaDecoderLayer(
+ (0-14): 15 x LlamaDecoderLayer(
(self_attn): LlamaSdpaAttention(
(q_proj): Linear(in_features=576, out_features=576, bias=False)
(k_proj): Linear(in_features=576, out_features=192, bias=False)
</details> <br/>
Train Dataset
Trained on 84,871,894 tokens from the wikimedia/wikipedia dataset.
- Num Samples:
99,800 - Subset:
20231101.en - Split:
train
Training Objective
DistillationObjective(
logits_loss_component=LossComponent(
weight=1,
loss_fn='kl'
),
hs_loss_component=LossComponent(
weight=0
),
attn_loss_component=LossComponent(
weight=0
)
)Hyperparameters
The following hyperparameters were used during training:
<details> <summary>Expand</summary>
- learning_rate:
0.0002 - trainbatchsize:
4 - evalbatchsize:
2 - seed:
42 - optimizer:
Adam with betas=(0.9,0.999) and epsilon=1e-08 - lrschedulertype:
polynomial - num_epochs:
1.0 - distillationobjective: `DistillationObjective( logitslosscomponent=LossComponent( weight=1, lossfn='kl' ), hslosscomponent=LossComponent( weight=0 ), attnlosscomponent=LossComponent( weight=0 ) )`
- lrscheduler: `<torch.optim.lrscheduler.LambdaLR object at 0x7eb253ff9660>`
- studentmodelnameorpath:
None - studentconfignameorpath:
None - studentmodelconfig:
{'num_hidden_layers': 15} - reinitialize_weights:
None - copyteachermodules:
[('lm_head', False)] - studentmodelas_bitnet:
False - studentmodeluse_liger:
False - teachermodelnameorpath:
HuggingFaceTB/SmolLM-135M - teacherloadin_8bit:
False - teacherloadin_4bit:
False - dataset_uri:
wikimedia/wikipedia - dataset_subset:
20231101.en - dataset_split:
train - datasetcolumnname:
text - datasetsamplesize:
100000 - datasettestsize:
0.002 - dataset_shuffle:
False - datasetshuffleseed:
42 - datasettrustremote_code:
False - gradientaccumulationsteps:
1 - weight_decay:
0.0 - maxgradnorm:
1.0 - warmup_ratio:
0.0 - warmup_steps:
0 - gradient_checkpointing:
True
</details> <br/>
Framework Versions
- Distily 0.5.0
- Transformers 4.44.2
- Pytorch 2.5.0.dev20240911+cu121
- Datasets 2.21.0
