fix: honor explicit model loader classes (#34880)
Co-authored-by: ehhuang <yinghai@meta.com>
This commit is contained in:
@@ -4278,6 +4278,9 @@ def get_model_loader(
|
|||||||
if load_config.load_format == LoadFormat.DUMMY:
|
if load_config.load_format == LoadFormat.DUMMY:
|
||||||
return DummyModelLoader(load_config)
|
return DummyModelLoader(load_config)
|
||||||
|
|
||||||
|
if isinstance(load_config.load_format, type):
|
||||||
|
return load_config.load_format(load_config)
|
||||||
|
|
||||||
if model_config and model_config.quantization in ["auto-round-int8"]:
|
if model_config and model_config.quantization in ["auto-round-int8"]:
|
||||||
logger.info("Using IncModelLoader due to AutoRound quantization config.")
|
logger.info("Using IncModelLoader due to AutoRound quantization config.")
|
||||||
return IncModelLoader(load_config)
|
return IncModelLoader(load_config)
|
||||||
@@ -4332,9 +4335,6 @@ def get_model_loader(
|
|||||||
)
|
)
|
||||||
return ModelOptModelLoader(load_config)
|
return ModelOptModelLoader(load_config)
|
||||||
|
|
||||||
if isinstance(load_config.load_format, type):
|
|
||||||
return load_config.load_format(load_config)
|
|
||||||
|
|
||||||
if load_config.load_format == LoadFormat.SHARDED_STATE:
|
if load_config.load_format == LoadFormat.SHARDED_STATE:
|
||||||
return ShardedStateLoader(load_config)
|
return ShardedStateLoader(load_config)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user