Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 34 additions & 2 deletions cosmos_framework/inference/common/args.py
Original file line number Diff line number Diff line change
Expand Up @@ -603,6 +603,38 @@ def build_checkpoint(self, *, checkpoints: dict[str, CheckpointConfig]) -> Check
CfgpSize = Annotated[int, pydantic.Field(ge=1, le=2)]
CompiledRegion = Literal["all", "language"]

# Low-precision quantization method to apply to the model at load time.
# One of ``mxfp8`` / ``nvfp4``, or ``None`` (default) to disable.
# Routed to the VFM model loader, which selects an FSDP-compatible
# (module-swap) path when sharded (``dp_shard_size > 1``) and an in-place
# path when replicated (``dp_shard_size == 1``). Note ``mxfp8`` / ``nvfp4``
# are only supported on the replicated path.
QuantizationMethod = Literal["mxfp8", "nvfp4"]


class QuantizationArgs(ArgsBase):
"""Low-precision quantization arguments applied to the model at load time."""

quantization_method: QuantizationMethod | None
quantization_include_regex: list[str]
quantization_exclude_regex: list[str]


class QuantizationOverrides(OverridesBase):
quantization_method: QuantizationMethod | None = None
"""Quantization method (``mxfp8`` / ``nvfp4``), or ``None`` to disable.

Post-training quantization (PTQ) is applied in-place to the model at load
time. Only supported on Blackwell architectures and when FSDP sharding is disabled.
"""
quantization_include_regex: list[str] = ["language_model.model.layers"]
"""Regexes matched against module FQNs; a Linear is quantized only if it matches one (empty = all)."""
quantization_exclude_regex: list[str] = pydantic.Field(default_factory=list)
"""Regexes matched against module FQNs; a Linear is skipped if it matches any."""

def build_quantization(self) -> QuantizationArgs:
return self._build(QuantizationArgs)


class ParallelismArgs(ArgsBase):
"""Parallelism arguments."""
Expand Down Expand Up @@ -702,7 +734,7 @@ class GuardrailOverrides(OverridesBase):
"""Offload guardrail models to CPU."""


class SetupArgs(ABC, CheckpointArgs, ParallelismArgs, GuardrailArgs):
class SetupArgs(ABC, CheckpointArgs, ParallelismArgs, QuantizationArgs, GuardrailArgs):
output_dir: ResolvedPath
keep_going: bool
skip_invalid_samples: bool
Expand Down Expand Up @@ -737,7 +769,7 @@ def get_variant(cls) -> str:
return cls.model_fields["variant"].default


class SetupOverrides(ABC, CheckpointOverrides, ParallelismOverrides, GuardrailOverrides):
class SetupOverrides(ABC, CheckpointOverrides, ParallelismOverrides, QuantizationOverrides, GuardrailOverrides):
"""Inference setup arguments."""

output_dir: Annotated[ResolvedPath | None, tyro.conf.arg(aliases=("-o",))] = None
Expand Down
1 change: 1 addition & 0 deletions cosmos_framework/inference/common/public_model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
"projects.cosmos3.vfm.configs.base.defaults.model_config.RectifiedFlowInferenceConfig": "rectified_flow_inference_config",
"projects.cosmos3.vfm.configs.base.defaults.model_config.RectifiedFlowTrainingConfig": "rectified_flow_training_config",
"projects.cosmos3.vfm.configs.base.defaults.parallelism.ParallelismConfig": "parallelism_config",
"projects.cosmos3.vfm.configs.base.defaults.quantization.QuantizationConfig": "quantization_config",
"projects.cosmos3.vfm.configs.base.defaults.vlm.PretrainedWeightsConfig": "pretrained_weights_config",
"projects.cosmos3.vfm.configs.base.defaults.vlm.VLMConfig": "vlm_config",
}
Expand Down
12 changes: 12 additions & 0 deletions cosmos_framework/inference/inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@

from cosmos_framework.configs.base.defaults.compile import CompileConfig
from cosmos_framework.configs.base.defaults.parallelism import ParallelismConfig
from cosmos_framework.configs.base.defaults.quantization import QuantizationConfig
from cosmos_framework.inference.args import (
ModelMode,
NegativeMetadataMode,
Expand Down Expand Up @@ -1055,6 +1056,14 @@ def _get_compile_config(cls, setup_args: ParallelismArgs) -> CompileConfig:
compile_dynamic=setup_args.compile_dynamic,
)

@classmethod
def _get_quantization_config(cls, setup_args: SetupArgs) -> QuantizationConfig:
return QuantizationConfig(
method=setup_args.quantization_method,
include_regex=list(setup_args.quantization_include_regex),
exclude_regex=list(setup_args.quantization_exclude_regex),
)

@override
@classmethod
def _create(cls, setup_args: SetupArgs, **kwargs: Any) -> Self:
Expand All @@ -1064,6 +1073,7 @@ def _create(cls, setup_args: SetupArgs, **kwargs: Any) -> Self:
sampler_override = setup_args.sampler
parallelism_config = cls._get_parallelism_config(setup_args)
compile_config = cls._get_compile_config(setup_args)
quantization_config = cls._get_quantization_config(setup_args)
if setup_args.checkpoint_type == CheckpointType.DCP and setup_args.config_file_type == ConfigFileType.MODULE:
from cosmos_framework.inference.common.config import save_config
from cosmos_framework.utils.generator.model_loader import load_model_from_checkpoint
Expand All @@ -1081,6 +1091,7 @@ def _create(cls, setup_args: SetupArgs, **kwargs: Any) -> Self:
credential_path=setup_args.credential_path or None,
parallelism_config=attrs.asdict(parallelism_config),
compile_config=attrs.asdict(compile_config),
quantization_config=attrs.asdict(quantization_config),
load_ema_to_reg=setup_args.use_ema_weights,
experiment_opts=[
*setup_args.experiment_overrides,
Expand Down Expand Up @@ -1130,6 +1141,7 @@ def _create(cls, setup_args: SetupArgs, **kwargs: Any) -> Self:
config=config,
parallelism_config=parallelism_config,
compile_config=compile_config,
quantization_config=quantization_config,
).model
if model.config.rectified_flow_inference_config.scheduler_type != sampler_override:
model.config.rectified_flow_inference_config.scheduler_type = sampler_override
Expand Down
15 changes: 15 additions & 0 deletions cosmos_framework/inference/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@

from cosmos_framework.configs.base.defaults.compile import CompileConfig
from cosmos_framework.configs.base.defaults.parallelism import ParallelismConfig
from cosmos_framework.configs.base.defaults.quantization import QuantizationConfig
from cosmos_framework.inference.common.args import CheckpointType
from cosmos_framework.inference.common.checkpoints import register_checkpoints
from cosmos_framework.inference.common.config import structure_config, undo_config_dict_replacements, unstructure_config
Expand Down Expand Up @@ -416,6 +417,16 @@ def compile(self, value: dict | None):
return
self.model.setdefault("config", {})["compile"] = unstructure_config(CompileConfig(**value))

@property
def quantization(self) -> dict:
return self.model.get("config", {}).get("quantization", {})

@quantization.setter
def quantization(self, value: dict | None):
if value is None:
return
self.model.setdefault("config", {})["quantization"] = unstructure_config(QuantizationConfig(**value))


class Cosmos3OmniModel(transformers.PreTrainedModel):
config_class = Cosmos3OmniConfig # type: ignore
Expand Down Expand Up @@ -448,15 +459,19 @@ def from_pretrained_dcp(
config: Cosmos3OmniConfig | None = None,
parallelism_config: ParallelismConfig | None = None,
compile_config: CompileConfig | None = None,
quantization_config: QuantizationConfig | None = None,
):
if config is None:
config = Cosmos3OmniConfig.from_pretrained(checkpoint_path)
if parallelism_config is None:
parallelism_config = ParallelismConfig()
if compile_config is None:
compile_config = CompileConfig()
if quantization_config is None:
quantization_config = QuantizationConfig()
config.parallelism = attrs.asdict(parallelism_config)
config.compile = attrs.asdict(compile_config)
config.quantization = attrs.asdict(quantization_config)
model = cls(config)
checkpoint_type = CheckpointType.from_path(checkpoint_path)
match checkpoint_type:
Expand Down