File size: 16,216 Bytes
f97c0fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1faa658
f97c0fc
 
1faa658
f21fc37
90226a6
f21fc37
1faa658
b30d489
1faa658
 
 
b30d489
1faa658
 
 
f21fc37
1faa658
 
 
 
 
 
 
 
 
 
 
 
 
 
f21fc37
1faa658
b30d489
1faa658
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f21fc37
 
 
 
1faa658
b30d489
1faa658
 
 
 
 
 
 
 
 
 
 
b7ca20b
a72f4a2
 
 
 
 
63ad778
b7ca20b
63ad778
a72f4a2
 
 
 
d2f0510
 
1f68979
 
 
 
a72f4a2
48151fe
 
 
b7ca20b
 
48151fe
 
 
b30d489
376881d
b30d489
 
 
b7ca20b
d645a66
 
 
 
6bc0c61
 
adc1fc4
 
 
 
 
 
 
 
 
b7ca20b
 
 
 
 
 
54d80c7
b30d489
 
 
 
1faa658
 
 
c18cf18
1faa658
 
 
b30d489
1faa658
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f21fc37
1faa658
 
 
 
f21fc37
1faa658
 
 
 
 
 
f21fc37
1faa658
f21fc37
1faa658
 
 
 
 
 
 
 
f21fc37
1faa658
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
---
language:
- tr
license: llama3.1
base_model: meta-llama/Llama-3.1-8B-Instruct
tags:
- legal
- turkish
- llama-3.1
- fp8
- bfloat16
- mixed-precision
- question-answering
- fsdp-v2
- pytorch
- distributed-training
datasets:
- newmindai/EuroHPC-Legal
library_name: transformers
pipeline_tag: text-generation
model-index:
- name: Llama-3.1-8B-Instruct_w16a8_rw_with_gw_hp
  results: []
---

## **Abstract**
The Llama-3.1-8B-Instruct_w16a8_rw_with_gw_hp model is a Turkish legal instruction-tuned variant of Llama-3.1-8B-Instruct, trained using a modified Float8 Rowwise recipe **Rowwise scaling with grad weight in High Precision** that preserves gradient-weight computations in high precision (BF16). This recipe was developed as part of the “FSDP2 with Float8 Precision for Faster Training” study to explore how selective high-precision treatment of specific tensor paths affects both training stability and overall efficiency.
## **Experiment Context**

This model was trained with the Float8 modified rowwise scaling method that preserves the gradient weight computation in high precision (bfloat16) to enhance accuracy recipe. In this recipe during the forward pass, both the input and weight granularities are set to AXISWISE with a default E4M3 data type, while in the backward pass for gradient input, the input granularity remains AXISWISE and the weight granularity is TENSORWISE, both using E4M3; finally, in the backward pass for gradient weight, scaling is disabled for both input and grad output—keeping them in high precision and E4M3 respectively—and the configuration further enables round_scales_to_power_of_2=True to ensure numerical stability. 
```python
from torchao.float8 import (
    convert_to_float8_training,
    Float8LinearConfig)
config = Float8LinearConfig.from_recipe_name("rowwise_with_gw_hp")
model = convert_to_float8_training(model, config=config)
````
## **Base Model Technical Specifications**
- **Parameters**: 8 Billion
- **Architecture Family**: Llama 3.1
- **Maximum Position Embeddings**: 131,072
- **Attention Heads**: 32 (`num_attention_heads`)
- **Key-Value Heads**: 8 (`num_key_value_heads`)
- **Hidden Layers**: 32 (`num_hidden_layers`)
- **Hidden Size**: 4,096 (`hidden_size`)
- **Intermediate Size**: 14,336
- **Vocabulary Size**: 128,256
- **Precision**: bfloat16
- **RoPE Scaling**: type `llama3`, factor = 8.0
- **RMS Norm Epsilon**: 1e-05
- **Activation**: SiLU

## **Training Methodology**

### *Training Configuration*
- **Model**: `meta-llama/Llama-3.1-8B-Instruct`
- **Sequence Length**: 4,096 (`seq_len`)
- **Epochs**: 1
- **Per-Device Micro Batch Size**: 2
- **Gradient Accumulation**: 4
- **GPUs**: 4 (via `CUDA_VISIBLE_DEVICES=0,1,2,3`)
- **dtype**: `bf16` && `fp8=false`
  - Weights: bfloat16
  - Activations: bfloat16
- **Optimizer**: AdamW
  - Learning Rate: 2e-5
  - Weight Decay: 0.01
  - Betas: (0.9, 0.95)
  - Epsilon: 1e-8
- **LR Scheduler**: Cosine; warmup = 10% (`warmup_ratio=0.1`) | also `warmup_steps=100`
- **Max Grad Norm**: 1.0
- **Gradient Checkpointing**: Enabled
- **Evaluation**: every 5 steps (`eval_steps=5`, `eval_samples=1000`)
- **Checkpointing**: every 10 steps; keep last 5; select best by `eval_loss`
- **Logging**: every step to file; Weights & Biases in offline mode
- **Seed**: 100
- **Distributed Training**: `torch.distributed.run` (single node, multi-GPU)
  - FSDP2 (Optimized Fully Sharded Data Parallel)
### *Setups*
- **Precision**: Used Half-precision bfloat16 as data type and for computation.
- **Hardware**: HPC (EuroHPC/BSC-class) node with 4 × NVIDIA H100 GPUs. 
- **Framework**: PyTorch with `torchrun` for distributed training. 

### *Dependencies*
| package | Version |
|-------|--------|
| Transformers  | 4.57.1 |
| torch  | 2.9.0+cu128 |
| accelerate  | 0.14.1 |
| datasets  | 4.3.0 |
| huggingface-hub  | 0.36.0 |
| tensorboard | 2.20.0 |
| tensorboard-data-server | 0.7.2 |
| wandb  | 0.22.1 |


## Job Details
| model                                    | Job ID   | Runtime (mins) | Nodes | GPUs | Node-hour | GPU-hour   | micro-batch | batch-size | gradient_accumulation | total_batch_size |
| ---------------------------------------- | -------- | -------------- | ----- | ---- | --------- | ---------- | ----------- | ---------- | --------------------- | ---------------- |
| Llama-3.1-8B-Instruct_w16a8_rw           | 31768103 | 115.75         | 1     | 4    | **1.929** |  **7.716** | 2           | 2          | 4                     | 32               |
| Llama-3.1-8B-Instruct_w16a8_rw_with_gw_hp| 31837629 | 109.00         | 1     | 4    | **1.816** |  **7.266** | 2           | 2          | 4                     | 32               |
| Llama-3.1-8B-Instruct-w16a8-mxtw         | 31768031 | 64.00          | 1     | 4    | **1.066** |  **4.266** | 2           | 2          | 4                     | 32               |
| Llama-3.1-8B-Instruct-w16a16-tw          | 31768074 | 138.75         | 1     | 4    | **2,312** | **9,25**   | 2           | 2          | 4                     | 32               |
| Llama-3.1-8B-Instruct-w16a8-1node-bs8    | 31768093 | 123.75         | 1     | 4    | **2.062** | **8,250**  | 2           | 2          | 4                     | 32               |
| Llama-3.1-8B-Instruct-w16a16-4nodes-bs32 | 31478433 | 31.75          | 4     | 4    | **2.117** | **8.467**  | 4           | 4          | 8                     | 512              |
| Llama-3.1-8B-Instruct-w16a8-4nodes-bs32  | 31478468 | 39.75          | 4     | 4    | **2.650** | **10.600** | 4           | 4          | 8                     | 512              |
| Llama-3.1-8B-Instruct-w16a16-8nodes-bs32 | 31476914 | 22.00          | 8     | 4    | **2.933** | **11.733** | 4           | 4          | 8                     | 1024             |
| Llama-3.1-8B-Instruct-w16a8-8nodes-bs32  | 31476844 | 23.50          | 8     | 4    | **3.133** | **12.533** | 4           | 4          | 8                     | 1024             |
| Llama-3.1-8B-Instruct-w16a16-8nodes-bs64 | 31476914 | 22.00          | 8     | 4    | **2.933** | **11.733** | 4           | 8          | 8                     | 1024             |
| Llama-3.1-8B-Instruct-w16a8-8nodes-bs64  | 31476844 | 23.50          | 8     | 4    | **3.133** | **12.533** | 4           | 8          | 8                     | 1024             |
| Llama-3.1-8B-Instruct-w16a8-rw_4nodes            | 33477070 | 39.75          | 4     | 4    | **2.650** | **10.600** | 4           | 4          | 8                     | 512              |
| Llama-3.1-8B-Instruct-w16a8-rw-8nodes            | 33476690 | 23.50          | 8     | 4    | **3.133** | **12.533** | 4           | 4          | 8                     | 1024             |
| Llama-3.1-8B-Instruct-w16a8-rw_with_gw_hp_4nodes | 33477179 | 37.43          | 4     | 4    | **2.495** | **9.982**  | 4           | 4          | 8                     | 512              |
| Llama-3.1-8B-Instruct-w16a8-rw-with-gw-hp-8nodes | 33476618 | 22.13          | 8     | 4    | **2.951** | **11.802** | 4           | 4          | 8                     | 1024             |

### *Training Time Analysision*
| Model                                               | Training Time (mins) | Memory Allocated (avg %) | GPU Utilization (avg %) | Speed vs bf16 |
| :-------------------------------------------------- | --------------------: | -----------------------: | -----------------------: | -------------: |
| **Llama-3.1-8B-Instruct_w16a16-tw**                 | 138.75267            | 74.4189                  | 56.6059%                 | _             |
| **Llama-3.1-8B-Instruct-w16a8-1node-bs8**           | 123.75267            | 68.8982                  | 97.5364%                 | 12.11%        |
| **Llama-3.1-8B-Instruct_w16a8_rw**                  | 115.75364            | 69.6132                  | 97.7689%                 | 19.87%        |
| **Llama-3.1-8B-Instruct_w16a8_rw_with_gw_hp**       | 109.00364            | 69.4806                  | 97.3312%                 | 27.33%        |
| **Llama-3.1-8B-Instruct-w16a8-mxtw**                | 64.00328             | 68.8982                  | 95.5661%                 | 116.82%       |
### *2-models trained on 1Node with fp8 recipes*
| Loss metric results for w16a16 & rowwise_with_gw_hp recipe| Memory allocation for w16a16 & rowwise_with_gw_hp recipe | Utilization for w16a16 & rowwise_with_gw_hp recipe | 
|---------|---------|---------|
| ![lossRWGWHP](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/VAPmLlCaZPaks9SCSiGnW.png) | ![memALRWGWHP](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/WHhSPl1n2BpDhqzGl_ljh.png) | ![gpuutilsRWGWHP](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/PdiR1e2SGyTOloURHy19G.png) |

# *All 15-models trained on(1Node,4Noes,8Nodes with both bfp16-fp8 && bfp16 configurations and fp8 recipes)*
| perplexity metric results for bfp16 && bfp16-fp8 configurations | Accuracy metric results for bfp16 && bfp16-fp8 configurations | Loss metric results for bfp16 && bfp16-fp8 configurations | Memory allocation for bfp16 && bfp16-fp8 configurations | Utilization for bfp16 && bfp16-fp8 configurations |
|:--:|:--:|:--:|:--:|:--:|
| ![perp](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/ij1hlr8E2qvdZM4uGC7lq.png) | ![acc](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/7lO8mVKPnQQkyUTw8H6GA.png) | ![train_loss](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/E73tvcC6u9VrvTIkznwU2.png) | ![memAlo](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/NsHL_yaTtnjwD1e4EHcLP.png) | ![utils](https://cdn-uploads.huggingface.co/production/uploads/683d4880e639f8d647355997/5mqF8xcRWuZdC_sGS9FCe.png) |

| Model                                           | Max Loss (train) | Min Loss (train) | Avg Loss (train) | Final Loss (train) | ± Std (train) | Max Loss (val) | Min Loss (val) | Avg Loss (val) | Final Loss (val) | ± Std (val) |
| ---------------------------------------------- | ---------------- | ---------------- | ---------------- | ------------------ | ------------- | -------------- | -------------- | -------------- | ---------------- | ----------- |
| Llama-3.1-8B-Instruct-w16a8-rw           | 8            | 3.1682           | 0.5740           | 0.8118           | 0.6431             | 0.2746        | 1.0613         | 0.8394         | 0.8937         | 0.8394           | 0.0688      |
| Llama-3.1-8B-Instruct_w16a8_rw_with_gw_hp| 8         | 3.1837           | 0.5763           | 0.8116           | 0.6420             | 0.2751        | 1.0599         | 0.8391         | 0.8933         | 0.8391           | 0.0685      |
| Llama-3.1-8B-Instruct-w16a8-mxtw         | 8          | 3.1983           | 0.5747           | 0.8115           | 0.6446             | 0.2758        | 1.0562         | 0.8384         | 0.8923         | 0.8384           | 0.0677      |
| Llama-3.1-8B-Instruct-w16a16-tw          | 8          | 3.1235           | 0.7203           | 0.9750           | 0.3344        | 0.7612             | 1.9113         | 0.8907         | 0.9831         | 0.1897      | 0.8907           | 312        | 
| Llama-3.1-8B-Instruct-w16a8-1node-bs8    | 8          | 3.1661           | 0.7261           | 0.9804           | 0.3374        | 0.7672             | 1.9230         | 0.8948         | 0.9867         | 0.1906      | 0.8951           | 312        | 
| Llama-3.1-8B-Instruct-w16a16-4nodes-bs32 | 32         | 3.2452           | 0.7414           | 0.9665           | 0.4844        | 0.7504             | 1.0538         | 0.8382         | 0.8844         | 0.0725      | 0.8382           | 70         | 
| Llama-3.1-8B-Instruct-w16a8-4nodes-bs32  | 32         | 3.2840           | 0.7478           | 0.9748           | 0.4905        | 0.7581             | 1.0701         | 0.8430         | 0.8922         | 0.0764      | 0.8430           | 70         | 
| Llama-3.1-8B-Instruct-w16a16-8nodes-bs32 | 32         | 3.2311           | 0.8448           | 1.1856           | 0.6434        | 0.8448             | 1.0257         | 0.8977         | 0.9460         | 0.0568      | 0.8977           | 35         | 
| Llama-3.1-8B-Instruct-w16a8-8nodes-bs32  | 32         | 3.3003           | 0.8473           | 1.1866           | 0.6481        | 0.8473             | 1.0203         | 0.8992         | 0.9445         | 0.0539      | 0.8992           | 35         | 
| Llama-3.1-8B-Instruct-w16a16-4nodes-bs64 | 64         | 3.2311           | 0.8448           | 1.1856           | 0.6434        | 0.8448             | 1.0257         | 0.8977         | 0.9460         | 0.0568      | 0.8977           | 35         | 
| Llama-3.1-8B-Instruct-w16a8-8nodes-bs64  | 64         | 3.3003           | 0.8473           | 1.1866           | 0.6481        | 0.8473             | 1.0203         | 0.8992         | 0.9445         | 0.0539      | 0.8992           | 17         |
| Llama-3.1-8B-Instruct-w16a8-rw_4nodes             | 64 | 3.4517          | 0.7624           | 1.1173           | 0.7624        | 0.6891             | 1.3225         | 0.8791         | 0.9732         | 0.8791      | 0.1612           | 35         | 
| Llama-3.1-8B-Instruct-w16a8-rw_8nodes             | 64 | 3.8944          | 0.9583           | 1.6423           | 0.9583        | 1.0117             | 1.5384         | 1.0253         | 1.2103         | 1.0253      | 0.2849           | 17         | 
| Llama-3.1-8B-Instruct-w16a8-rw_with_gw_hp_4nodes  | 64 | 3.4517          | 0.7481           | 1.1091           | 0.7481        | 0.7021             | 1.3393         | 0.8660         | 0.9641         | 0.8666      | 0.1732           | 35         |
| Llama-3.1-8B-Instruct-w16a8-rw_with_gw_hp_8nodes  | 64 | 3.9289          | 0.9702           | 1.6514           | 0.9702        | 1.0127             | 1.5537         | 1.0377         | 1.2222         | 1.0377      | 0.2877           | 17         |

## **Implementation**
### *Gpu && Memory usage Profiling*
The training progress has been profiled using **pytorch-profiler** tool. 
- follow the steps to visualize the profiles: 
  1. pip install the versions that mentioned in the dependencies section of these libs tensorboard and tensorboard-data-server.
  2. Visualize pytorch profiles by runing the command provided below. 
  ```python
  tensorboard --logdir="./Llama-3.1-8B-Instruct_w16a8_rowwise_with_gw_hp" --port="6006"
  ````


### *Usage*

**Note**: the final model has been saved in bfloat16 format. For inference, load the model in bfloat16 or float16 as shown below:

```python
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
model_name = "newmindai/Llama-3.1-8B-Instruct_w16a8_rw_with_gw_hp"
dtype = torch.bfloat16
tok = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(
    model_name,
    torch_dtype=dtype,
    device_map="auto"
)
prompt = "Soru: Kişisel Verilerin Korunması Kanunu uyarınca hangi durumlarda açık rıza aranmaz? Cevap:"
inputs = tok(prompt, return_tensors="pt").to(model.device)
with torch.no_grad():
    out = model.generate(
        **inputs,
        max_new_tokens=256,
        do_sample=False
    )

print(tok.decode(out[0], skip_special_tokens=True))
````
## **Ethical Considerations and Disclaimers**
* Research & development purposes only; not a substitute for professional legal counsel.
* Users must ensure compliance with data protection and sector regulations.
* Potential biases may exist in domain data and model outputs.

## **Model & Data Card Metadata**

* **Total Parameters**: 8,030,261,248
* **Serialized Size (approx.)**: 16,060,522,496 bytes
* **Config precision**: bfloat16
* **RoPE**: llama3 scaling, factor 8.0

## **References and Citations**

### *Base Model*
```bibtex
@misc{meta_llama31_8b_instruct,
  title={Llama 3.1 8B Instruct},
  author={Meta AI},
  year={2024},
  howpublished={\url{https://huggingface.co/meta-llama/Llama-3.1-8B-Instruct}}
}
```
### *Training Dataset*

```bibtex
@misc{euro_hpc_legal,
  title={EuroHPC-Legal},
  author={newmindai},
  year={2025},
  howpublished={\url{https://huggingface.co/datasets/newmindai/EuroHPC-Legal}}
}
```