google / google/saxml

saxml/tools/convert_llama_ckpt.py casts weights to float16, losing precision

Open
#28 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Python
Stars
156
Forks
35
PR merge metrics
No merged PRs in 30d

Description

The weights for e.g. Meta-Llama-3.1-70B-Instruct are distributed in bfloat16 format. When converting the weights, the saxml script first casts the weights to float16, which is lossy.

E.g. for Meta-Llama-3.1-70B-Instruct:
```
>>> example = torch.load('consolidated.01.pth', weights_only=True, map_location=torch.device('cpu'), mmap=True)['layers.79.feed_forward.w1.weight'][100][5685]
>>> example
tensor(-4.2617e-06, dtype=torch.bfloat16)
>>> example.type(torch.float16)
tensor(-4.2915e-06, dtype=torch.float16)
```

(This sounds similar to an issue HuggingFace had with weight conversion: https://github.com/huggingface/transformers/issues/25446, which was acknowledged to degrade performance and was fixed.)

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.