diff --git a/flytekit/core/node.py b/flytekit/core/node.py index f579d391ad..0b3b24369e 100644 --- a/flytekit/core/node.py +++ b/flytekit/core/node.py @@ -232,7 +232,19 @@ def with_overrides( logger.warning("This override is beta. We may want to revisit this in the future.") if not isinstance(task_config, type(self.run_entity._task_config)): raise ValueError("can't change the type of the task config") + + # Add debugging for task config overrides + old_config = self.run_entity._task_config + logger.info(f"Applying task_config override on node {self.id}") + logger.info(f" Old config: {old_config}") + logger.info(f" New config: {task_config}") + + # For Elastic PyTorch tasks, specifically log nnodes changes + if hasattr(old_config, 'nnodes') and hasattr(task_config, 'nnodes'): + logger.info(f" Elastic PyTorch nnodes override: {old_config.nnodes} -> {task_config.nnodes}") + self.run_entity._task_config = task_config + logger.info(f"Task config override applied successfully") if container_image is not None: assert_not_promise(container_image, "container_image") diff --git a/flytekit/tools/translator.py b/flytekit/tools/translator.py index 60c1ef11c7..3ac9eb79bd 100644 --- a/flytekit/tools/translator.py +++ b/flytekit/tools/translator.py @@ -465,7 +465,44 @@ def get_serializable_node( override_pod_spec = _serialize_pod_spec( entity._pod_template, entity.flyte_entity._get_container(settings), settings ) - task_spec = get_serializable(entity_mapping, settings, entity.flyte_entity, options=options) + + # Fix 2: Override-aware serialization + # Create a temporary copy of the task with overrides applied for serialization + # This ensures that get_custom() sees the correct configuration without affecting the original task + from flytekit import logger + import copy + + task_entity = entity.flyte_entity + + # Check if we need to create a temporary task copy with overrides applied + # This happens when the node's task config differs from the original task config + needs_override_copy = False + if hasattr(entity, 'run_entity') and hasattr(entity.run_entity, '_task_config'): + node_task_config = entity.run_entity._task_config + original_task_config = task_entity._task_config + + # Compare configs to see if they differ (indicating an override was applied) + if node_task_config != original_task_config: + needs_override_copy = True + logger.info(f"Fix 2: Node {entity.id} has task config override, creating temporary copy for serialization") + logger.info(f" Original task config: {original_task_config}") + logger.info(f" Node task config: {node_task_config}") + + # For Elastic PyTorch tasks, specifically log nnodes changes + if hasattr(original_task_config, 'nnodes') and hasattr(node_task_config, 'nnodes'): + logger.info(f" Elastic PyTorch nnodes override: {original_task_config.nnodes} -> {node_task_config.nnodes}") + + if needs_override_copy: + # Create a shallow copy of the task entity and apply the node's task config + # This ensures get_custom() sees the overridden config during serialization + temp_task_entity = copy.copy(task_entity) + temp_task_entity._task_config = node_task_config + + logger.info(f"Fix 2: Using temporary task copy with overridden config for serialization") + task_spec = get_serializable(entity_mapping, settings, temp_task_entity, options=options) + else: + # No override needed, use the original task entity + task_spec = get_serializable(entity_mapping, settings, task_entity, options=options) node_model = workflow_model.Node( id=_dnsify(entity.id), metadata=entity.metadata, diff --git a/plugins/flytekit-kf-pytorch/flytekitplugins/kfpytorch/task.py b/plugins/flytekit-kf-pytorch/flytekitplugins/kfpytorch/task.py index 1972f10bd9..9d5f5e925d 100644 --- a/plugins/flytekit-kf-pytorch/flytekitplugins/kfpytorch/task.py +++ b/plugins/flytekit-kf-pytorch/flytekitplugins/kfpytorch/task.py @@ -507,13 +507,21 @@ def execute(self, **kwargs) -> Any: return self._execute(**kwargs) def get_custom(self, settings: SerializationSettings) -> Optional[Dict[str, Any]]: - if self._task_config.nnodes == 1: + # Get the current task config, which should include any overrides applied via with_overrides() + current_nnodes = self._task_config.nnodes + + # Add debugging to help diagnose override issues + logger.info(f"PytorchElasticFunctionTask.get_custom() called with nnodes={current_nnodes} (type: {type(current_nnodes)})") + + if current_nnodes == 1: """ Torch elastic distributed training is executed in a normal k8s pod so that this works without the kubeflow train operator. """ + logger.info("Using single-node execution (normal Kubernetes pod)") return super().get_custom(settings) else: + logger.info(f"Using multi-node execution (PyTorchJob) with nnodes={current_nnodes}") from flyteidl.plugins.kubeflow.pytorch_pb2 import ElasticConfig try: @@ -521,7 +529,8 @@ def get_custom(self, settings: SerializationSettings) -> Optional[Dict[str, Any] except ImportError: raise ImportError(TORCH_IMPORT_ERROR_MESSAGE) - min_nodes, max_nodes = run.parse_min_max_nnodes(str(self._task_config.nnodes)) + min_nodes, max_nodes = run.parse_min_max_nnodes(str(current_nnodes)) + logger.info(f"Parsed nnodes '{current_nnodes}' to min_nodes={min_nodes}, max_nodes={max_nodes}") elastic_config = ElasticConfig( rdzv_backend=self.rdzv_backend, diff --git a/plugins/flytekit-kf-pytorch/tests/test_elastic_task.py b/plugins/flytekit-kf-pytorch/tests/test_elastic_task.py index f8742d1fe9..e69eae1d4b 100644 --- a/plugins/flytekit-kf-pytorch/tests/test_elastic_task.py +++ b/plugins/flytekit-kf-pytorch/tests/test_elastic_task.py @@ -310,3 +310,119 @@ def test_task(): test_task() assert e.value.timestamp is not None + + +def test_nnodes_override() -> None: + """Test that nnodes overrides work correctly in get_custom().""" + from flytekit import workflow + from flytekit.configuration import SerializationSettings + + # Create a task with default nnodes=2 (multi-node) + @task(task_config=Elastic(nnodes=2, nproc_per_node=1)) + def multi_node_task(): + return "done" + + @workflow + def test_workflow(): + # Override to single-node + return multi_node_task().with_overrides( + task_config=Elastic(nnodes=1, nproc_per_node=1) + ) + + # Get the workflow node + wf = test_workflow + node = list(wf.nodes)[0] + + # Verify the override was applied to the task config + assert node.run_entity._task_config.nnodes == 1 + + # Test serialization - this is where the bug would manifest + settings = SerializationSettings(image_config=None) + + # Original task should return custom config (PyTorchJob for multi-node) + original_custom = multi_node_task.get_custom(settings) + assert original_custom is not None, "Multi-node task should have custom config" + + # Overridden task should return None (normal pod for single-node) + overridden_custom = node.run_entity.get_custom(settings) + assert overridden_custom is None, "Single-node override should have no custom config" + + # Verify they're different (override took effect) + assert original_custom != overridden_custom, "Override should change the custom config" + + +def test_nnodes_override_reverse() -> None: + """Test overriding from single-node to multi-node.""" + from flytekit import workflow + from flytekit.configuration import SerializationSettings + + # Create a task with default nnodes=1 (single-node) + @task(task_config=Elastic(nnodes=1, nproc_per_node=1)) + def single_node_task(): + return "done" + + @workflow + def test_workflow(): + # Override to multi-node + return single_node_task().with_overrides( + task_config=Elastic(nnodes=2, nproc_per_node=1) + ) + + # Get the workflow node + wf = test_workflow + node = list(wf.nodes)[0] + + # Verify the override was applied to the task config + assert node.run_entity._task_config.nnodes == 2 + + # Test serialization + settings = SerializationSettings(image_config=None) + + # Original task should return None (normal pod for single-node) + original_custom = single_node_task.get_custom(settings) + assert original_custom is None, "Single-node task should have no custom config" + + # Overridden task should return custom config (PyTorchJob for multi-node) + overridden_custom = node.run_entity.get_custom(settings) + assert overridden_custom is not None, "Multi-node override should have custom config" + + # Verify they're different (override took effect) + assert original_custom != overridden_custom, "Override should change the custom config" + + +def test_fix2_serialization_isolation() -> None: + """Test Fix 2: Ensure serialization with overrides doesn't affect original task.""" + from flytekit import workflow + from flytekit.configuration import SerializationSettings + from flytekit.tools.translator import get_serializable_workflow + from collections import OrderedDict + + # Create a task with default nnodes=2 (multi-node) + @task(task_config=Elastic(nnodes=2, nproc_per_node=1)) + def multi_node_task(): + return "done" + + @workflow + def test_workflow(): + # Override to single-node + return multi_node_task().with_overrides( + task_config=Elastic(nnodes=1, nproc_per_node=1) + ) + + # Test serialization with Fix 2 + settings = SerializationSettings(image_config=None) + entity_mapping = OrderedDict() + + # Serialize the workflow - this should trigger Fix 2 logic + wf_spec = get_serializable_workflow(entity_mapping, settings, test_workflow) + + # Verify that the original task config was not permanently modified + assert multi_node_task._task_config.nnodes == 2, "Original task config should remain unchanged" + + # Verify that the node still has the override applied + node = list(test_workflow.nodes)[0] + assert node.run_entity._task_config.nnodes == 1, "Node should still have override applied" + + # Verify that serialization produced the correct result + # The serialized workflow should contain a task with single-node config + assert wf_spec is not None, "Workflow serialization should succeed" diff --git a/test_elastic_override_fix.py b/test_elastic_override_fix.py new file mode 100644 index 0000000000..4d67666e26 --- /dev/null +++ b/test_elastic_override_fix.py @@ -0,0 +1,83 @@ +#!/usr/bin/env python3 +""" +Test to verify that the Elastic PyTorch task override fix works correctly. +This test demonstrates the issue and verifies the fix. +""" + +import logging +from flytekitplugins.kfpytorch.task import Elastic +from flytekit import task, workflow +from flytekit.configuration import SerializationSettings + +# Set up logging to see our debug messages +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + + +@task(task_config=Elastic(nnodes=2, nproc_per_node=1)) +def train_task(): + """A training task with default nnodes=2 (multi-node).""" + print("Training...") + return "done" + + +@workflow +def training_workflow(): + """Workflow that overrides the task to use nnodes=1 (single-node).""" + # This should override the task to use single-node execution + single_node_task = train_task().with_overrides( + task_config=Elastic(nnodes=1, nproc_per_node=1) + ) + return single_node_task + + +def test_override_fix(): + """Test that demonstrates the fix for Elastic PyTorch task overrides.""" + print("=== Testing Elastic PyTorch Task Override Fix ===\n") + + # Get the workflow and its node + wf = training_workflow + node = list(wf.nodes)[0] # First (and only) node + + print(f"Original task nnodes: {train_task._task_config.nnodes}") + print(f"Overridden node task nnodes: {node.run_entity._task_config.nnodes}") + + # Test serialization - this is where the bug manifests + settings = SerializationSettings(image_config=None) + + print("\n=== Testing Serialization ===") + print("Getting custom config for original task:") + original_custom = train_task.get_custom(settings) + + print("\nGetting custom config for overridden task:") + overridden_custom = node.run_entity.get_custom(settings) + + print(f"\nOriginal custom config: {original_custom}") + print(f"Overridden custom config: {overridden_custom}") + + # Analyze the results + print("\n=== Results Analysis ===") + + # Original task should use PyTorchJob (multi-node) + if original_custom is None: + print("✅ Original task correctly uses single-node (None custom config)") + else: + print("✅ Original task correctly uses multi-node (PyTorchJob custom config)") + + # Overridden task should use normal pod (single-node) + if overridden_custom is None: + print("✅ SUCCESS: Overridden task correctly uses single-node execution!") + print(" The override from nnodes=2 to nnodes=1 is working properly.") + else: + print("❌ ISSUE: Overridden task still uses multi-node execution") + print(" The override is not taking effect properly.") + + # Additional verification + if original_custom != overridden_custom: + print("✅ Override is taking effect - configurations are different") + else: + print("❌ Override not working - configurations are identical") + + +if __name__ == "__main__": + test_override_fix() diff --git a/test_elastic_override_issue.py b/test_elastic_override_issue.py new file mode 100644 index 0000000000..f3f28f9acc --- /dev/null +++ b/test_elastic_override_issue.py @@ -0,0 +1,57 @@ +#!/usr/bin/env python3 +""" +Test script to reproduce the Elastic PyTorch task override issue. +""" + +from flytekitplugins.kfpytorch.task import Elastic +from flytekit import task, workflow +from flytekit.configuration import SerializationSettings + + +# Create a task with default nnodes=2 +@task(task_config=Elastic(nnodes=2, nproc_per_node=1)) +def train_task(): + print("Training...") + return "done" + + +@workflow +def training_workflow(): + # Override the task to use nnodes=1 + single_node_task = train_task().with_overrides(task_config=Elastic(nnodes=1, nproc_per_node=1)) + return single_node_task + + +def test_serialization(): + """Test what happens during serialization.""" + print("=== Original Task Config ===") + print(f"Original task nnodes: {train_task._task_config.nnodes}") + + # Get the workflow + wf = training_workflow + + # Get the node from the workflow + node = list(wf.nodes)[0] # First node + print(f"Node task config nnodes: {node.run_entity._task_config.nnodes}") + + # Test serialization + settings = SerializationSettings(image_config=None) + + print("\n=== Serialization Test ===") + print("Original task get_custom:") + original_custom = train_task.get_custom(settings) + print(f"Original custom config: {original_custom}") + + print("\nOverridden task get_custom:") + overridden_custom = node.run_entity.get_custom(settings) + print(f"Overridden custom config: {overridden_custom}") + + # Check if the issue exists + if original_custom == overridden_custom: + print("\n❌ ISSUE CONFIRMED: Override not taking effect in serialization!") + else: + print("\n✅ Override working correctly in serialization") + + +if __name__ == "__main__": + test_serialization()