huggingface / huggingface/diffusers

Loading pipeline in precision it was saved in

オープン
#9,797 コメント 10 件 リアクション 3 件 担当者 0 名 GitHub で見る
enhancement
主要言語
Python
スター
34.5k
フォーク
7.3k
平均マージ
3日 3時間
マージ済み PR(30日)
91

説明

**Is your feature request related to a problem? Please describe.**
Currently, if `torch_dtype` is not specified, the pipeline defaults to loading in `float32`. This behavior causes `float16` or `bfloat16` weights to be upcast to `float32` when the model is saved in lower precision, leading to increased memory usage. In scenarios where memory efficiency is critical (e.g., when exporting the model to another format), it’s important to load the model in the original precision specified in the safetensors file. Additionally, there’s currently no way to determine the dtype the model was saved in.

**Describe the solution you'd like.**
A feature similar to `torch_dtype="auto"` in the transformers library would be helpful. This option allows models to be loaded with the dtype defined in their configuration. However, diffuser pipeline models generally lack a dtype specification in their configs. It is sometimes possible to use `torch_dtype` from `text_encoder` config, but not all pipelines have it and it is not clear if this is a reliable place to check the precision of the model.

**Describe alternatives you've considered.**
A possible solution could be implementing a method to identify the model’s precision prior to calling `from_pretrained`, as the weights are accessible only after the model is downloaded inside `from_pretrained` and remain hidden from external access. This approach would allow users to set the appropriate `torch_dtype` for loading the model.

**Additional context.**
This feature is relevant to `optimum-cli` use cases where model conversion or export to other formats must work within memory constraints. If there’s already a way to achieve this, guidance would be appreciated.

コントリビューションガイド

コントリビューションガイドを開く

調査の方向性

まず、from_pretrained を通じたパイプラインの読み込みを追跡し、読み込み中に safetensors の重みがどのように利用可能になるかを調べます。要求されている動作を transformers の torch_dtype="auto" および記載されている optimum-cli の変換ユースケースと比較します。ユーザーが torch_dtype を手動で指定しなくても、保存された精度を確実に検出または保持できる方法があれば完了です。

索引モデルが issue の本文から書いたものです。

評価

技術スタック
python, pytorch
領域
machine-learning, performance
issue の種類
機能追加
難易度
5/5
見積もり時間
1週間以上
活発さ
停滞
明瞭さ
説明が足りない
初心者へのやさしさ
35/100

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。