[diffusion] chore: minor cleanups (#19123)
This commit is contained in:
@@ -244,7 +244,7 @@ class MinimalA2AAttnOp(DistributedAttention):
|
|||||||
SparseLinearAttentionBackend,
|
SparseLinearAttentionBackend,
|
||||||
SageSparseLinearAttentionBackend,
|
SageSparseLinearAttentionBackend,
|
||||||
):
|
):
|
||||||
logger.warning(
|
logger.warning_once(
|
||||||
"TurboWan now only supports `sla_attn` or `sage_sla_attn` and has been automatically set to attention_type. Please set --attention-backend to `sla_attn` or `sage_sla_attn`."
|
"TurboWan now only supports `sla_attn` or `sage_sla_attn` and has been automatically set to attention_type. Please set --attention-backend to `sla_attn` or `sage_sla_attn`."
|
||||||
)
|
)
|
||||||
if attention_type == "sagesla":
|
if attention_type == "sagesla":
|
||||||
|
|||||||
+6
-9
@@ -1,19 +1,17 @@
|
|||||||
import json
|
|
||||||
import os
|
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.models import ModelConfig
|
from sglang.multimodal_gen.configs.models import ModelConfig
|
||||||
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
|
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
|
||||||
TextEncoderLoader,
|
TextEncoderLoader,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import _clean_hf_config_inplace
|
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
||||||
|
get_diffusers_component_config,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class ImageEncoderLoader(TextEncoderLoader):
|
class ImageEncoderLoader(TextEncoderLoader):
|
||||||
|
|
||||||
component_names = ["image_encoder"]
|
component_names = ["image_encoder"]
|
||||||
expected_library = "transformers"
|
expected_library = "transformers"
|
||||||
|
|
||||||
@@ -41,10 +39,9 @@ class ImageEncoderLoader(TextEncoderLoader):
|
|||||||
# revision=server_args.revision,
|
# revision=server_args.revision,
|
||||||
# model_override_args=None,
|
# model_override_args=None,
|
||||||
# )
|
# )
|
||||||
with open(os.path.join(component_model_path, "config.json")) as f:
|
model_config = get_diffusers_component_config(
|
||||||
model_config = json.load(f)
|
component_path=component_model_path
|
||||||
_clean_hf_config_inplace(model_config)
|
)
|
||||||
logger.debug("HF model config: %s", model_config)
|
|
||||||
|
|
||||||
encoder_config = server_args.pipeline_config.image_encoder_config
|
encoder_config = server_args.pipeline_config.image_encoder_config
|
||||||
encoder_config.update_model_arch(model_config)
|
encoder_config.update_model_arch(model_config)
|
||||||
|
|||||||
@@ -21,7 +21,6 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.fsdp_load import shard_model
|
from sglang.multimodal_gen.runtime.loader.fsdp_load import shard_model
|
||||||
from sglang.multimodal_gen.runtime.loader.utils import (
|
from sglang.multimodal_gen.runtime.loader.utils import (
|
||||||
_clean_hf_config_inplace,
|
|
||||||
set_default_torch_dtype,
|
set_default_torch_dtype,
|
||||||
skip_init_modules,
|
skip_init_modules,
|
||||||
)
|
)
|
||||||
@@ -173,20 +172,12 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
self, component_model_path: str, server_args: ServerArgs, component_name: str
|
||||||
):
|
):
|
||||||
"""Load the text encoders based on the model path, and inference args."""
|
"""Load the text encoders based on the model path, and inference args."""
|
||||||
# model_config: PretrainedConfig = get_hf_config(
|
|
||||||
# model=model_path,
|
|
||||||
# trust_remote_code=server_args.trust_remote_code,
|
|
||||||
# revision=server_args.revision,
|
|
||||||
# model_override_args=None,
|
|
||||||
# )
|
|
||||||
diffusers_pretrained_config = get_config(
|
diffusers_pretrained_config = get_config(
|
||||||
component_model_path, trust_remote_code=True
|
component_model_path, trust_remote_code=True
|
||||||
)
|
)
|
||||||
model_config = get_diffusers_component_config(
|
model_config = get_diffusers_component_config(
|
||||||
component_path=component_model_path
|
component_path=component_model_path
|
||||||
)
|
)
|
||||||
_clean_hf_config_inplace(model_config)
|
|
||||||
logger.debug("HF model config: %s", model_config)
|
|
||||||
|
|
||||||
def is_not_first_encoder(module_name):
|
def is_not_first_encoder(module_name):
|
||||||
return "2" in module_name
|
return "2" in module_name
|
||||||
|
|||||||
@@ -135,22 +135,21 @@ class GPUWorker:
|
|||||||
# otherwise empty offloaded weights could fail lora converting
|
# otherwise empty offloaded weights could fail lora converting
|
||||||
if self.server_args.dit_layerwise_offload:
|
if self.server_args.dit_layerwise_offload:
|
||||||
# enable layerwise offload if possible
|
# enable layerwise offload if possible
|
||||||
for dit in filter(
|
for module_name in [
|
||||||
None,
|
"transformer",
|
||||||
[
|
"transformer_2",
|
||||||
self.pipeline.get_module("transformer"),
|
"video_dit",
|
||||||
self.pipeline.get_module("transformer_2"),
|
"video_dit_2",
|
||||||
self.pipeline.get_module("video_dit"),
|
"audio_dit",
|
||||||
self.pipeline.get_module("video_dit_2"),
|
]:
|
||||||
self.pipeline.get_module("audio_dit"),
|
dit = self.pipeline.get_module(module_name)
|
||||||
],
|
if dit:
|
||||||
):
|
if isinstance(dit, OffloadableDiTMixin):
|
||||||
if isinstance(dit, OffloadableDiTMixin):
|
dit.configure_layerwise_offload(self.server_args)
|
||||||
dit.configure_layerwise_offload(self.server_args)
|
else:
|
||||||
else:
|
logger.info(
|
||||||
logger.info(
|
f"Module {type(dit).__name__} does not support layerwise offload. Skipping."
|
||||||
f"Module {type(dit).__name__} does not support layerwise offload. Skipping."
|
)
|
||||||
)
|
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Worker {self.rank}: Initialized device, model, and distributed environment."
|
f"Worker {self.rank}: Initialized device, model, and distributed environment."
|
||||||
@@ -199,7 +198,7 @@ class GPUWorker:
|
|||||||
logger.info(
|
logger.info(
|
||||||
f"Peak GPU memory: {peak_reserved_gb:.2f} GB, "
|
f"Peak GPU memory: {peak_reserved_gb:.2f} GB, "
|
||||||
f"Peak allocated: {peak_allocated_gb:.2f} GB, "
|
f"Peak allocated: {peak_allocated_gb:.2f} GB, "
|
||||||
f"Memory pool overhead: {pool_overhead_gb:.2f} GB ({pool_overhead_gb/peak_reserved_gb*100:.1f}%), "
|
f"Memory pool overhead: {pool_overhead_gb:.2f} GB ({pool_overhead_gb / peak_reserved_gb * 100:.1f}%), "
|
||||||
f"Remaining GPU memory at peak: {remaining_gpu_mem_gb:.2f} GB. "
|
f"Remaining GPU memory at peak: {remaining_gpu_mem_gb:.2f} GB. "
|
||||||
f"Components that could stay resident (based on the last request workload): {can_stay_resident}. "
|
f"Components that could stay resident (based on the last request workload): {can_stay_resident}. "
|
||||||
f"Related offload server args to disable: {suggested_args_str}"
|
f"Related offload server args to disable: {suggested_args_str}"
|
||||||
|
|||||||
@@ -40,12 +40,12 @@ from requests.exceptions import ConnectionError as RequestsConnectionError
|
|||||||
from requests.exceptions import RequestException
|
from requests.exceptions import RequestException
|
||||||
from safetensors import safe_open
|
from safetensors import safe_open
|
||||||
from transformers import AutoConfig, PretrainedConfig
|
from transformers import AutoConfig, PretrainedConfig
|
||||||
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
|
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.layers.quantization import (
|
from sglang.multimodal_gen.runtime.layers.quantization import (
|
||||||
QuantizationConfig,
|
QuantizationConfig,
|
||||||
get_quantization_config,
|
get_quantization_config,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.utils import _clean_hf_config_inplace
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock
|
from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
@@ -209,6 +209,34 @@ def _ci_validate_diffusers_model(model_path: str) -> tuple[bool, bool]:
|
|||||||
return True, False
|
return True, False
|
||||||
|
|
||||||
|
|
||||||
|
def _verify_diffusers_model_complete(path: str) -> bool:
|
||||||
|
"""Check if a diffusers model directory has all required component subdirectories."""
|
||||||
|
config_path = os.path.join(path, "model_index.json")
|
||||||
|
if not os.path.exists(config_path):
|
||||||
|
return False
|
||||||
|
|
||||||
|
try:
|
||||||
|
with open(config_path) as config_file:
|
||||||
|
model_index = json.load(config_file)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Failed to read model_index.json at %s: %s", config_path, exc)
|
||||||
|
return False
|
||||||
|
|
||||||
|
component_keys = [
|
||||||
|
key
|
||||||
|
for key, value in model_index.items()
|
||||||
|
if isinstance(value, (list, tuple))
|
||||||
|
and len(value) == 2
|
||||||
|
and all(isinstance(item, str) for item in value)
|
||||||
|
]
|
||||||
|
if component_keys:
|
||||||
|
return all(os.path.exists(os.path.join(path, key)) for key in component_keys)
|
||||||
|
|
||||||
|
return os.path.exists(os.path.join(path, "transformer")) and os.path.exists(
|
||||||
|
os.path.join(path, "vae")
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
_CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = {
|
_CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = {
|
||||||
# ChatGLMConfig.model_type: ChatGLMConfig,
|
# ChatGLMConfig.model_type: ChatGLMConfig,
|
||||||
# DbrxConfig.model_type: DbrxConfig,
|
# DbrxConfig.model_type: DbrxConfig,
|
||||||
@@ -235,8 +263,7 @@ def get_hf_config(
|
|||||||
model_override_args: dict | None = None,
|
model_override_args: dict | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> PretrainedConfig:
|
) -> PretrainedConfig:
|
||||||
is_gguf = check_gguf_file(component_model_path)
|
if check_gguf_file(component_model_path):
|
||||||
if is_gguf:
|
|
||||||
raise NotImplementedError("GGUF models are not supported.")
|
raise NotImplementedError("GGUF models are not supported.")
|
||||||
|
|
||||||
config = AutoConfig.from_pretrained(
|
config = AutoConfig.from_pretrained(
|
||||||
@@ -253,13 +280,6 @@ def get_hf_config(
|
|||||||
if model_override_args:
|
if model_override_args:
|
||||||
config.update(model_override_args)
|
config.update(model_override_args)
|
||||||
|
|
||||||
# Special architecture mapping check for GGUF models
|
|
||||||
if is_gguf:
|
|
||||||
if config.model_type not in MODEL_FOR_CAUSAL_LM_MAPPING_NAMES:
|
|
||||||
raise RuntimeError(f"Can't get gguf config for {config.model_type}.")
|
|
||||||
model_type = MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type]
|
|
||||||
config.update({"architectures": [model_type]})
|
|
||||||
|
|
||||||
return config
|
return config
|
||||||
|
|
||||||
|
|
||||||
@@ -270,14 +290,9 @@ def get_config(
|
|||||||
model_override_args: Optional[dict] = None,
|
model_override_args: Optional[dict] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
try:
|
return AutoConfig.from_pretrained(
|
||||||
config = AutoConfig.from_pretrained(
|
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
|
||||||
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
|
)
|
||||||
)
|
|
||||||
except ValueError as e:
|
|
||||||
raise e
|
|
||||||
|
|
||||||
return config
|
|
||||||
|
|
||||||
|
|
||||||
def load_dict(file_path):
|
def load_dict(file_path):
|
||||||
@@ -305,7 +320,6 @@ def get_diffusers_component_config(
|
|||||||
if not os.path.exists(component_path):
|
if not os.path.exists(component_path):
|
||||||
component_path = maybe_download_model(component_path)
|
component_path = maybe_download_model(component_path)
|
||||||
|
|
||||||
# tokenizer
|
|
||||||
config_names = ["generation_config.json"]
|
config_names = ["generation_config.json"]
|
||||||
# By default, we load config.json, but scheduler_config.json for scheduler
|
# By default, we load config.json, but scheduler_config.json for scheduler
|
||||||
if "scheduler" in component_path:
|
if "scheduler" in component_path:
|
||||||
@@ -321,6 +335,10 @@ def get_diffusers_component_config(
|
|||||||
lambda acc, path: acc | load_dict(path), config_file_paths, {}
|
lambda acc, path: acc | load_dict(path), config_file_paths, {}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
_clean_hf_config_inplace(combined_config)
|
||||||
|
|
||||||
|
logger.debug("HF model config: %s", combined_config)
|
||||||
|
|
||||||
return combined_config
|
return combined_config
|
||||||
|
|
||||||
|
|
||||||
@@ -414,9 +432,9 @@ CONTEXT_LENGTH_KEYS = [
|
|||||||
def attach_additional_stop_token_ids(tokenizer):
|
def attach_additional_stop_token_ids(tokenizer):
|
||||||
# Special handling for stop token <|eom_id|> generated by llama 3 tool use.
|
# Special handling for stop token <|eom_id|> generated by llama 3 tool use.
|
||||||
if "<|eom_id|>" in tokenizer.get_added_vocab():
|
if "<|eom_id|>" in tokenizer.get_added_vocab():
|
||||||
tokenizer.additional_stop_token_ids = set(
|
tokenizer.additional_stop_token_ids = {
|
||||||
[tokenizer.get_added_vocab()["<|eom_id|>"]]
|
tokenizer.get_added_vocab()["<|eom_id|>"]
|
||||||
)
|
}
|
||||||
else:
|
else:
|
||||||
tokenizer.additional_stop_token_ids = None
|
tokenizer.additional_stop_token_ids = None
|
||||||
|
|
||||||
@@ -635,45 +653,11 @@ def maybe_download_model(
|
|||||||
Local path to the model
|
Local path to the model
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def _verify_diffusers_model_complete(path: str) -> bool:
|
|
||||||
"""Check if model directory (of a diffusers model, not a component) has required subdirectories."""
|
|
||||||
config_path = os.path.join(path, "model_index.json")
|
|
||||||
if not os.path.exists(config_path):
|
|
||||||
return False
|
|
||||||
|
|
||||||
try:
|
|
||||||
with open(config_path) as config_file:
|
|
||||||
model_index = json.load(config_file)
|
|
||||||
except Exception as exc:
|
|
||||||
logger.warning(
|
|
||||||
"Failed to read model_index.json at %s: %s", config_path, exc
|
|
||||||
)
|
|
||||||
return False
|
|
||||||
|
|
||||||
component_keys = [
|
|
||||||
key
|
|
||||||
for key, value in model_index.items()
|
|
||||||
if isinstance(value, (list, tuple))
|
|
||||||
and len(value) == 2
|
|
||||||
and all(isinstance(item, str) for item in value)
|
|
||||||
]
|
|
||||||
if component_keys:
|
|
||||||
return all(
|
|
||||||
os.path.exists(os.path.join(path, component_key))
|
|
||||||
for component_key in component_keys
|
|
||||||
)
|
|
||||||
|
|
||||||
transformer_dir = os.path.join(path, "transformer")
|
|
||||||
vae_dir = os.path.join(path, "vae")
|
|
||||||
return os.path.exists(transformer_dir) and os.path.exists(vae_dir)
|
|
||||||
|
|
||||||
# 1. Local path check: if path exists locally, verify it's complete (skip for LoRA)
|
# 1. Local path check: if path exists locally, verify it's complete (skip for LoRA)
|
||||||
if os.path.exists(model_name_or_path):
|
if os.path.exists(model_name_or_path):
|
||||||
# TODO: lots of duplication here
|
|
||||||
if not force_diffusers_model:
|
if not force_diffusers_model:
|
||||||
return model_name_or_path
|
return model_name_or_path
|
||||||
elif is_lora or _verify_diffusers_model_complete(model_name_or_path):
|
if is_lora or _verify_diffusers_model_complete(model_name_or_path):
|
||||||
# CI validation: check all subdirectories for missing shards
|
|
||||||
if not is_lora:
|
if not is_lora:
|
||||||
is_valid, cleanup_performed = _ci_validate_diffusers_model(
|
is_valid, cleanup_performed = _ci_validate_diffusers_model(
|
||||||
model_name_or_path
|
model_name_or_path
|
||||||
@@ -687,8 +671,6 @@ def maybe_download_model(
|
|||||||
)
|
)
|
||||||
# Fall through to download
|
# Fall through to download
|
||||||
else:
|
else:
|
||||||
# Local path is not in HF cache structure, can't clean up
|
|
||||||
# Raise error since we can't fix this automatically
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"CI validation failed for local model at {model_name_or_path}. "
|
f"CI validation failed for local model at {model_name_or_path}. "
|
||||||
"Some safetensors shards are missing. "
|
"Some safetensors shards are missing. "
|
||||||
@@ -722,25 +704,21 @@ def maybe_download_model(
|
|||||||
)
|
)
|
||||||
if not force_diffusers_model:
|
if not force_diffusers_model:
|
||||||
return str(local_path)
|
return str(local_path)
|
||||||
elif is_lora or _verify_diffusers_model_complete(local_path):
|
if is_lora or _verify_diffusers_model_complete(local_path):
|
||||||
# CI validation: check all subdirectories for missing shards
|
|
||||||
if not is_lora:
|
if not is_lora:
|
||||||
is_valid, cleanup_performed = _ci_validate_diffusers_model(local_path)
|
is_valid, cleanup_performed = _ci_validate_diffusers_model(local_path)
|
||||||
if not is_valid:
|
if not is_valid:
|
||||||
if cleanup_performed:
|
logger.warning(
|
||||||
logger.warning(
|
"CI validation failed for cached model at %s, "
|
||||||
"CI validation failed for cached model at %s, "
|
"%s, will re-download",
|
||||||
"cache has been cleaned up, will re-download",
|
local_path,
|
||||||
local_path,
|
(
|
||||||
)
|
"cache has been cleaned up"
|
||||||
# Fall through to download
|
if cleanup_performed
|
||||||
else:
|
else "cleanup was not performed"
|
||||||
# This shouldn't happen for HF cache paths, but handle it
|
),
|
||||||
logger.warning(
|
)
|
||||||
"CI validation failed for cached model at %s, "
|
# Fall through to download
|
||||||
"but cleanup was not performed, will attempt re-download",
|
|
||||||
local_path,
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
logger.info("Found complete model in cache at %s", local_path)
|
logger.info("Found complete model in cache at %s", local_path)
|
||||||
return str(local_path)
|
return str(local_path)
|
||||||
@@ -748,7 +726,6 @@ def maybe_download_model(
|
|||||||
logger.info("Found complete model in cache at %s", local_path)
|
logger.info("Found complete model in cache at %s", local_path)
|
||||||
return str(local_path)
|
return str(local_path)
|
||||||
else:
|
else:
|
||||||
# Model found in cache but incomplete
|
|
||||||
if not download:
|
if not download:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Model {model_name_or_path} found in cache but is incomplete and download=False."
|
f"Model {model_name_or_path} found in cache but is incomplete and download=False."
|
||||||
@@ -930,3 +907,54 @@ def get_metadata_from_safetensors_file(file_path: str):
|
|||||||
return metadata
|
return metadata
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(e)
|
logger.warning(e)
|
||||||
|
|
||||||
|
|
||||||
|
def get_quant_config_from_safetensors_metadata(
|
||||||
|
file_path: str,
|
||||||
|
) -> Optional[QuantizationConfig]:
|
||||||
|
"""Extract quantization config from a safetensors file's metadata header.
|
||||||
|
|
||||||
|
Safetensors files can embed a flat string→string metadata dict in their header.
|
||||||
|
We expect a ``quantization_config`` key containing a JSON-encoded dict with at
|
||||||
|
least a ``quant_method`` field (e.g. ``"fp8"``), matching the format written by
|
||||||
|
``convert_hf_to_fp8.py`` when embedded into a config.json.
|
||||||
|
|
||||||
|
Returns None if no recognizable quantization metadata is found.
|
||||||
|
"""
|
||||||
|
metadata = get_metadata_from_safetensors_file(file_path)
|
||||||
|
if not metadata:
|
||||||
|
return None
|
||||||
|
|
||||||
|
quant_config_str = metadata.get("quantization_config")
|
||||||
|
if not quant_config_str:
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
quant_config_dict = json.loads(quant_config_str)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"failed to parse quantization_config from safetensors metadata: %s", e
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
quant_method = quant_config_dict.get("quant_method")
|
||||||
|
if not quant_method:
|
||||||
|
logger.warning(
|
||||||
|
"quantization_config in safetensors metadata is missing 'quant_method'"
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
quant_cls = get_quantization_config(quant_method)
|
||||||
|
config = quant_cls.from_config(quant_config_dict)
|
||||||
|
logger.info(
|
||||||
|
"loaded quantization config (%s) from safetensors metadata: %s",
|
||||||
|
quant_method,
|
||||||
|
file_path,
|
||||||
|
)
|
||||||
|
return config
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
"failed to build QuantizationConfig from safetensors metadata: %s", e
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|||||||
@@ -115,7 +115,7 @@ class ConversionResult:
|
|||||||
with self.lock:
|
with self.lock:
|
||||||
for k, v in q_weights.items():
|
for k, v in q_weights.items():
|
||||||
self.weight_map[k] = filename
|
self.weight_map[k] = filename
|
||||||
self.param_count += len(v)
|
self.param_count += v.numel()
|
||||||
self.modules_to_not_convert.extend(module_names)
|
self.modules_to_not_convert.extend(module_names)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user