"""Utilities adapted from* https://github.com/huggingface/transformers/blob/main/src/transformers/quantizers/quantizer_bnb_4bit.py* https://github.com/huggingface/transformers/blob/main/src/transformers/integrations/bitsandbytes.py"""importtorchimportbitsandbytesasbnbfromtransformers.quantizers.quantizers_utilsimportget_module_from_nameimporttorch.nnasnnfromaccelerateimportinit_empty_weightsdef_replace_with_bnb_linear(
model,
method="nf4",
has_been_replaced=False,
):
""" Private method that wraps the recursion for module replacement. Returns the converted model and a boolean that indicates if the conversion has been successfull or not. """forname, moduleinmodel.named_children():
ifisinstance(module, nn.Linear):
withinit_empty_weights():
in_features=module.in_featuresout_features=module.out_featuresifmethod=="llm_int8":
model._modules[name] =bnb.nn.Linear8bitLt(
in_features,
out_features,
module.biasisnotNone,
has_fp16_weights=False,
threshold=6.0,
)
has_been_replaced=Trueelse:
model._modules[name] =bnb.nn.Linear4bit(
in_features,
out_features,
module.biasisnotNone,
compute_dtype=torch.bfloat16,
compress_statistics=False,
quant_type="nf4",
)
has_been_replaced=True# Store the module class in case we need to transpose the weight latermodel._modules[name].source_cls=type(module)
# Force requires grad to False to avoid unexpected errorsmodel._modules[name].requires_grad_(False)
iflen(list(module.children())) >0:
_, has_been_replaced=_replace_with_bnb_linear(
module,
has_been_replaced=has_been_replaced,
)
# Remove the last key for recursionreturnmodel, has_been_replaceddefcheck_quantized_param(
model,
param_name: str,
) ->bool:
module, tensor_name=get_module_from_name(model, param_name)
ifisinstance(module._parameters.get(tensor_name, None), bnb.nn.Params4bit):
# Add here check for loaded components' dtypes once serialization is implementedreturnTrueelifisinstance(module, bnb.nn.Linear4bit) andtensor_name=="bias":
# bias could be loaded by regular set_module_tensor_to_device() from accelerate,# but it would wrongly use uninitialized weight there.returnTrueelse:
returnFalsedefcreate_quantized_param(
model,
param_value: "torch.Tensor",
param_name: str,
target_device: "torch.device",
state_dict=None,
unexpected_keys=None,
pre_quantized=False
):
module, tensor_name=get_module_from_name(model, param_name)
iftensor_namenotinmodule._parameters:
raiseValueError(f"{module} does not have a parameter or a buffer named {tensor_name}.")
old_value=getattr(module, tensor_name)
iftensor_name=="bias":
ifparam_valueisNone:
new_value=old_value.to(target_device)
else:
new_value=param_value.to(target_device)
new_value=torch.nn.Parameter(new_value, requires_grad=old_value.requires_grad)
module._parameters[tensor_name] =new_valuereturnifnotisinstance(module._parameters[tensor_name], bnb.nn.Params4bit):
raiseValueError("this function only loads `Linear4bit components`")
if (
old_value.device==torch.device("meta")
andtarget_devicenotin ["meta", torch.device("meta")]
andparam_valueisNone
):
raiseValueError(f"{tensor_name} is on the meta device, we need a `value` to put in on {target_device}.")
ifpre_quantized:
if (param_name+".quant_state.bitsandbytes__fp4"notinstate_dict) and (
param_name+".quant_state.bitsandbytes__nf4"notinstate_dict
):
raiseValueError(
f"Supplied state dict for {param_name} does not contain `bitsandbytes__*` and possibly other `quantized_stats` components."
)
quantized_stats= {}
fork, vinstate_dict.items():
# `startswith` to counter for edge cases where `param_name`# substring can be present in multiple places in the `state_dict`ifparam_name+"."inkandk.startswith(param_name):
quantized_stats[k] =vifunexpected_keysisnotNoneandkinunexpected_keys:
unexpected_keys.remove(k)
new_value=bnb.nn.Params4bit.from_prequantized(
data=param_value,
quantized_stats=quantized_stats,
requires_grad=False,
device=target_device,
)
else:
new_value=param_value.to("cpu")
kwargs=old_value.__dict__new_value=bnb.nn.Params4bit(new_value, requires_grad=False, **kwargs).to(target_device)
module._parameters[tensor_name] =new_value
@SunMarc
Since the Flux params are quite huge (if we include the text encoder2, autoencoder, and the diffusion model itself) -- it totals to more than 30GB.
https://huggingface.co/lllyasviel/flux1-dev-bnb-nf4 ships a single safetensors file that has the diffusion model in NF4.
Now, I was able to get this converted and load it into our
FluxTransformer2DModel,but I am not seeing any size (state dict size) benefits. I am seeing the size benefits (yay!).But loading seems to be not working yet. What am I missing? Will appreciate feedback.Here is a detailed rundown of what I have done so far.
convert_nf4_flux.py
generate.py
The image generates just fine.
But not sure why we're not seeing any size benefit here.But the loading seems broken (generated image is noise). Advise?I have uploaded the NF4 serialized state dict here: https://huggingface.co/sayakpaul/flux.1-dev-nf4Loading script is below:
load_from_nf4_and_generate.py
NF4 serialization and loading is working fine!