From ae67b71c649d10eadb9aa1402a92667518813e0c Mon Sep 17 00:00:00 2001 From: FunJim Date: Wed, 23 Sep 2026 15:16:54 +0800 Subject: [PATCH] fix: verify complete Megatron patches before updating installations --- .../tools/apply_megatron_patch.py | 214 ++++++++++--- tests/test_apply_megatron_patch.py | 280 ++++++++++++++++++ tests/test_glm5_hybrid.py | 11 +- 3 files changed, 464 insertions(+), 41 deletions(-) create mode 100644 tests/test_apply_megatron_patch.py diff --git a/src/mcore_bridge/tools/apply_megatron_patch.py b/src/mcore_bridge/tools/apply_megatron_patch.py index 317557f..94f8d6d 100644 --- a/src/mcore_bridge/tools/apply_megatron_patch.py +++ b/src/mcore_bridge/tools/apply_megatron_patch.py @@ -15,49 +15,191 @@ with `index_n_heads=32` the un-chunked fp32 `[seqlen_q, batch, heads, seqlen_k]` tensor is 8 GiB at sequence 8192 and 128 GiB when packed to 32768 -- and is bit-identical to the un-chunked version. -Idempotent. `git apply --3way` is used inside a git checkout, so an upstream edit outside the lines -we change does not block it and a real overlap is left as a visible conflict rather than dropped; -`patch(1)` is used for a pip-installed megatron, which is not a git tree. +Prepare an unused environment; do not apply from training ranks. Source checkouts retain +Git three-way merge support. Wheels only receive runtime files, not Megatron's unit tests. +Application is prepared and checked in a temporary directory before any target is written; +conflicts leave the checkout, including its index, unchanged. --check only verifies the patch. """ +import argparse import importlib.util +import os import pathlib +import re +import shutil import subprocess import sys +import tempfile +from typing import Dict, Iterator, List, Optional PATCH = pathlib.Path(__file__).resolve().parent.parent / 'patches' / 'megatron_glm53_dev.patch' -MARKER_FILE = 'megatron/core/transformer/transformer_config.py' -MARKER_SYMBOL = 'kda_two_stage_gates' -INSTALL_HINT = 'pip install -U git+https://github.com/NVIDIA/Megatron-LM.git@dev' - - -def megatron_root(): - """The directory the patch's `megatron/core/...` paths are relative to.""" - spec = importlib.util.find_spec('megatron.core') - if spec is None or not spec.origin: - sys.exit(f'megatron is not installed; install it first:\n {INSTALL_HINT}') - return pathlib.Path(spec.origin).resolve().parents[2] # /megatron/core/__init__.py - - -def is_applied(root): - target = pathlib.Path(root) / MARKER_FILE - return target.is_file() and MARKER_SYMBOL in target.read_text(errors='ignore') - - -def main(): - root = megatron_root() - if is_applied(root): - print(f'already applied: {root}') - return 0 - in_git = not subprocess.run(['git', 'rev-parse', '--show-toplevel'], cwd=root, capture_output=True).returncode - command = (['git', 'apply', '--3way', str(PATCH)] if in_git else ['patch', '-p1', '--forward', '-i', str(PATCH)]) - proc = subprocess.run(command, cwd=root, capture_output=True, text=True) - print(proc.stdout + proc.stderr, end='') - if proc.returncode: - print( - f'could not apply {PATCH.name} to {root}. If megatron drifted from the commit above, ' - f'update it and retry:\n {INSTALL_HINT}', - file=sys.stderr) - return proc.returncode + + +def megatron_root() -> pathlib.Path: + """Locate the package without importing megatron.core or initializing CUDA dependencies.""" + spec = importlib.util.find_spec('megatron') + if spec is not None and spec.submodule_search_locations: + for location in spec.submodule_search_locations: + if (pathlib.Path(location) / 'core' / '__init__.py').is_file(): + return pathlib.Path(location).resolve().parent + raise RuntimeError('Megatron is not installed; install the supported Megatron version first.') + + +def git_root(root: pathlib.Path) -> bool: + """A wheel inside another repository is not a Megatron source checkout.""" + if shutil.which('git') is None: + return False + result = subprocess.run(['git', 'rev-parse', '--show-toplevel'], cwd=root, capture_output=True, text=True) + return result.returncode == 0 and pathlib.Path(result.stdout.strip()).resolve() == root + + +def patch_files(source_checkout: bool) -> Dict[str, str]: + """Select existing-file diffs, preserving test changes only for source checkouts.""" + entries = {} + for section in re.split(r'(?=^diff --git )', PATCH.read_text(), flags=re.MULTILINE): + if not section.strip(): + continue + match = re.match(r'diff --git a/(\S+) b/\1\n', section) + if match is None: + raise RuntimeError('Unsupported patch header') + name = match[1] + if (not name.startswith(('megatron/core/', 'tests/')) or '..' in pathlib.PurePosixPath(name).parts + or name in entries or f'\n--- a/{name}\n+++ b/{name}\n' not in section): + raise RuntimeError(f'Unsupported patch target: {name}') + if source_checkout or name.startswith('megatron/core/'): + entries[name] = section + if not entries: + raise RuntimeError('Empty runtime patch') + return entries + + +def run_patch(root: pathlib.Path, + text: str, + reverse: bool = False, + dry_run: bool = True) -> subprocess.CompletedProcess: + command = ['patch', '--batch', '--force', '--fuzz=0', '--no-backup-if-mismatch', '-p1'] + command += ['--reverse'] if reverse else ['--forward'] + if dry_run: + command.append('--dry-run') + return subprocess.run(command, cwd=root, input=text, capture_output=True, text=True) + + +def patch_hunks(entries: Dict[str, str]) -> Iterator[str]: + """Check each hunk so a partially applied single file cannot pass as a fresh base.""" + for section in entries.values(): + parts = re.split(r'(?=^@@ )', section, flags=re.MULTILINE) + if len(parts) < 2: + raise RuntimeError('Patch target has no hunks') + for hunk in parts[1:]: + yield parts[0] + hunk + + +def prepare_patch(root: pathlib.Path, snapshot: pathlib.Path, entries: Dict[str, str], source_checkout: bool) -> None: + """Apply off-target, allowing a clean Git three-way merge when context has drifted.""" + text = ''.join(entries.values()) + merged = False + result = run_patch(snapshot, text) + if result.returncode == 0: + result = run_patch(snapshot, text, dry_run=False) + elif source_checkout: + merged = True + # Use a private index and worktree; borrow only objects needed for the merge base. + common = subprocess.run(['git', 'rev-parse', '--git-common-dir'], + cwd=root, + capture_output=True, + text=True, + check=True) + objects = (root / common.stdout.strip() / 'objects').resolve() + subprocess.run(['git', 'init', '-q', str(snapshot)], check=True) + subprocess.run(['git', 'add', '--', *entries], cwd=snapshot, check=True) + env = dict(os.environ, GIT_ALTERNATE_OBJECT_DIRECTORIES=str(objects)) + result = subprocess.run(['git', 'apply', '--3way', '--whitespace=nowarn', '-'], + cwd=snapshot, + input=text, + capture_output=True, + text=True, + env=env) + if result.returncode: + raise RuntimeError(f'Patch conflicts; no target files changed.\n{result.stdout}{result.stderr}') + # Git verifies a three-way result itself; its merged context may differ from the diff. + if not merged and run_patch(snapshot, text, reverse=True).returncode: + raise RuntimeError('Incomplete patch result; no target files changed.') + + +def write_files(root: pathlib.Path, snapshot: pathlib.Path, originals: Dict[str, bytes]) -> None: + """Restore original contents if publishing a prepared patch fails.""" + for name, content in originals.items(): + target = root / name + if target.resolve() != target or target.read_bytes() != content: + raise RuntimeError(f'Target changed during patch preparation: {name}') + written = [] + try: + for name in originals: + written.append(name) + (root / name).write_bytes((snapshot / name).read_bytes()) + except BaseException: + for name in written: + (root / name).write_bytes(originals[name]) + raise + + +def apply_patch(root: pathlib.Path, check_only: bool = False) -> None: + root = root.resolve() + source_checkout = git_root(root) + entries = patch_files(source_checkout) + originals = {} + for name in entries: + target = root / name + if not target.is_file() or target.resolve() != target: + raise RuntimeError(f'Missing or symlinked patch target: {target}') + originals[name] = target.read_bytes() + with tempfile.TemporaryDirectory(prefix='megatron-patch-') as directory: + snapshot = pathlib.Path(directory) + for name, content in originals.items(): + target = snapshot / name + target.parent.mkdir(parents=True, exist_ok=True) + target.write_bytes(content) + if run_patch(snapshot, ''.join(entries.values()), reverse=True).returncode == 0: + print(f'already applied (all {len(entries)} files verified): {root}') + return + partial = any(run_patch(snapshot, hunk, reverse=True).returncode == 0 for hunk in patch_hunks(entries)) + if not source_checkout: + if check_only: + raise RuntimeError(f'Patch not fully applied: {root}') + if partial: + raise RuntimeError(f'Partially applied patch; no target files changed: {root}') + prepare_patch(root, snapshot, entries, source_checkout) + # A three-way application to an already merged source tree is a no-op, even when + # unrelated edits changed the patch context and reverse dry-run could not match it. + if all((snapshot / name).read_bytes() == content for name, content in originals.items()): + print(f'already applied (all {len(entries)} files verified): {root}') + return + if check_only: + raise RuntimeError(f'Patch not fully applied: {root}') + if partial: + raise RuntimeError(f'Partially applied patch; no target files changed: {root}') + write_files(root, snapshot, originals) + print(f'applied and verified all {len(entries)} files: {root}') + + +def is_applied(root: pathlib.Path) -> bool: + try: + apply_patch(pathlib.Path(root), check_only=True) + except (RuntimeError, OSError, subprocess.CalledProcessError): + return False + return True + + +def main(argv: Optional[List[str]] = None) -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--root', type=pathlib.Path, help='Megatron source root; defaults to installed package') + parser.add_argument('--check', action='store_true', help='Verify all selected patch hunks without writing') + args = parser.parse_args(argv) + try: + apply_patch(args.root if args.root is not None else megatron_root(), args.check) + except (RuntimeError, OSError, subprocess.CalledProcessError) as error: + print(str(error), file=sys.stderr) + return 1 + return 0 if __name__ == '__main__': diff --git a/tests/test_apply_megatron_patch.py b/tests/test_apply_megatron_patch.py new file mode 100644 index 0000000..dea2845 --- /dev/null +++ b/tests/test_apply_megatron_patch.py @@ -0,0 +1,280 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Installer regressions using real patch/git executables, without Torch or CUDA.""" +import difflib +import hashlib +import importlib.util +import os +import pathlib +import subprocess +import tempfile +import unittest +import zipfile +from types import SimpleNamespace +from unittest.mock import patch + +TOOL = pathlib.Path(__file__).resolve().parents[1] / 'src/mcore_bridge/tools/apply_megatron_patch.py' +spec = importlib.util.spec_from_file_location('apply_megatron_patch', TOOL) +installer = importlib.util.module_from_spec(spec) +spec.loader.exec_module(installer) + + +def blob(data: str) -> str: + content = data.encode() + return hashlib.sha1(b'blob ' + str(len(content)).encode() + b'\0' + content).hexdigest() + + +def make_diff(name: str, before: str, after: str) -> str: + body = ''.join( + difflib.unified_diff(before.splitlines(True), after.splitlines(True), fromfile=f'a/{name}', tofile=f'b/{name}')) + return f'diff --git a/{name} b/{name}\nindex {blob(before)}..{blob(after)} 100644\n{body}' + + +class ApplyPatchTest(unittest.TestCase): + + def setUp(self): + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.root = pathlib.Path(temporary.name).resolve() + self.names = ['megatron/core/transformer_config.py', 'megatron/core/optimizer.py'] + self.before = 'before_context = 1\nvalue = 1\ncontext_a = 1\ncontext_b = 1\nafter_context = 1\n' + self.after = self.before.replace('value = 1', 'value = 2') + self.test_name = 'tests/unit_tests/test_example.py' + for name in self.names: + target = self.root / name + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(self.before) + self.patch_file = self.root / 'fixture.patch' + self.patch_file.write_text(''.join(make_diff(name, self.before, self.after) + for name in self.names) + make_diff(self.test_name, self.before, self.after)) + mock = patch.object(installer, 'PATCH', self.patch_file) + mock.start() + self.addCleanup(mock.stop) + + def git(self, *args): + return subprocess.run(['git', *args], cwd=self.root, check=True, capture_output=True).stdout + + def init_git(self): + target = self.root / self.test_name + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(self.before) + self.git('init', '-q') + self.git('add', '--', *self.names, self.test_name) + self.git('-c', 'user.name=Test', '-c', 'user.email=test@example.com', 'commit', '-qm', 'base') + + def snapshot(self): + return {str(p.relative_to(self.root)): p.read_bytes() for p in self.root.rglob('*') if p.is_file()} + + def test_wheel_apply_check_repeat_without_test_tree(self): + self.assertFalse(installer.is_applied(self.root)) + installer.apply_patch(self.root) + for name in self.names: + self.assertEqual((self.root / name).read_text(), self.after) + self.assertFalse((self.root / 'tests').exists()) + expected = self.snapshot() + installer.apply_patch(self.root, check_only=True) + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_marker_alone_is_not_a_complete_installation(self): + (self.root / self.names[0]).write_text('kda_two_stage_gates = False\n') + (self.root / self.names[1]).unlink() + expected = self.snapshot() + self.assertFalse(installer.is_applied(self.root)) + with self.assertRaisesRegex(RuntimeError, 'Missing'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_partial_patch_rejected(self): + (self.root / self.names[0]).write_text(self.after) + expected = self.snapshot() + with self.assertRaisesRegex(RuntimeError, 'Partially'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_partial_hunks_in_one_file_rejected(self): + before = self.before + '\n' * 10 + 'second = 1\n' + after = self.after + '\n' * 10 + 'second = 2\n' + self.patch_file.write_text(make_diff(self.names[0], before, after)) + (self.root / self.names[0]).write_text(before.replace('value = 1', 'value = 2')) + expected = self.snapshot() + with self.assertRaisesRegex(RuntimeError, 'Partially'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_partial_source_rejected_without_index_changes(self): + self.init_git() + (self.root / self.names[0]).write_text(self.after) + expected = self.snapshot() + with self.assertRaisesRegex(RuntimeError, 'Partially'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_missing_source_test_file_rejected(self): + self.init_git() + (self.root / self.test_name).unlink() + expected = self.snapshot() + with self.assertRaisesRegex(RuntimeError, 'Missing'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_failed_application_never_writes_target(self): + expected = self.snapshot() + original = installer.run_patch + + def fail_apply(root, text, reverse=False, dry_run=True): + if not dry_run: + (root / self.names[0]).write_text('partially applied') + return subprocess.CompletedProcess(['patch'], 1, '', 'simulated failure') + return original(root, text, reverse, dry_run) + + with patch.object(installer, 'run_patch', fail_apply): + with self.assertRaisesRegex(RuntimeError, 'simulated failure'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_unrelated_wheel_edits_preserved(self): + for name in self.names: + with (self.root / name).open('a') as target: + target.write('\n# unrelated upstream addition\n') + installer.apply_patch(self.root) + self.assertTrue(installer.is_applied(self.root)) + for name in self.names: + self.assertEqual((self.root / name).read_text(), self.after + '\n# unrelated upstream addition\n') + + def test_conflict_leaves_no_partial_files_or_rejects(self): + (self.root / self.names[1]).write_text('conflicting content\n') + expected = self.snapshot() + with self.assertRaisesRegex(RuntimeError, 'conflicts'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_symlink_rejected(self): + (self.root / self.names[0]).unlink() + (self.root / self.names[0]).symlink_to(self.root / self.names[1]) + with self.assertRaisesRegex(RuntimeError, 'symlinked'): + installer.apply_patch(self.root) + + def test_failed_publish_restores_originals(self): + expected = self.snapshot() + original = pathlib.Path.write_bytes + failed = False + + def fail_once(path, data): + nonlocal failed + if path == self.root / self.names[1] and not failed: + failed = True + original(path, b'partial write') + raise OSError('simulated write failure') + return original(path, data) + + with patch.object(pathlib.Path, 'write_bytes', fail_once): + with self.assertRaisesRegex(OSError, 'simulated'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_source_applies_test_diff_without_staging(self): + self.init_git() + index = (self.root / '.git/index').read_bytes() + installer.apply_patch(self.root) + for name in self.names + [self.test_name]: + self.assertEqual((self.root / name).read_text(), self.after) + self.assertEqual(index, (self.root / '.git/index').read_bytes()) + self.assertEqual(self.git('diff', '--cached'), b'') + self.assertTrue(installer.is_applied(self.root)) + + def test_three_way_preserves_context_changes_and_index(self): + self.init_git() + # Change a context line: patch --fuzz=0 fails, while a three-way merge is clean. + target = self.root / self.names[0] + target.write_text(self.before.replace('after_context = 1', 'after_context = 9')) + self.git('add', '--', self.names[0]) + index = (self.root / '.git/index').read_bytes() + installer.apply_patch(self.root) + self.assertEqual(target.read_text(), self.after.replace('after_context = 1', 'after_context = 9')) + self.assertEqual(index, (self.root / '.git/index').read_bytes()) + expected = self.snapshot() + installer.apply_patch(self.root, check_only=True) + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_three_way_conflict_preserves_worktree_and_index(self): + self.init_git() + (self.root / self.names[0]).write_text(self.before.replace('value = 1', 'value = 99')) + expected = self.snapshot() + with self.assertRaisesRegex(RuntimeError, 'conflicts'): + installer.apply_patch(self.root) + self.assertEqual(expected, self.snapshot()) + + def test_wheel_nested_in_git_does_not_require_tests(self): + self.init_git() + site = self.root / 'venv/lib/site-packages' + for name in self.names: + target = site / name + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(self.before) + self.assertFalse(installer.git_root(site)) + installer.apply_patch(site) + self.assertFalse((site / 'tests').exists()) + + def test_find_package_without_importing_core(self): + location = self.root / 'megatron' + (location / 'core/__init__.py').write_text('raise RuntimeError("must not import")\n') + with patch.object( + installer.importlib.util, + 'find_spec', + return_value=SimpleNamespace(submodule_search_locations=[str(location)])) as find: + self.assertEqual(installer.megatron_root(), self.root) + find.assert_called_once_with('megatron') + + def test_readonly_check_returns_nonzero_for_unpatched_root(self): + expected = self.snapshot() + self.assertEqual(installer.main(['--root', str(self.root), '--check']), 1) + self.assertEqual(expected, self.snapshot()) + + def test_packaged_runtime_patch_parses(self): + with patch.object(installer, 'PATCH', TOOL.parent.parent / 'patches/megatron_glm53_dev.patch'): + runtime = installer.patch_files(False) + source = installer.patch_files(True) + self.assertTrue(all(name.startswith('megatron/core/') for name in runtime)) + self.assertIn('megatron/core/transformer/transformer_config.py', runtime) + self.assertGreater(len(source), len(runtime)) + + +class PackagedPatchIntegrationTest(unittest.TestCase): + """Optionally exercise the entire bundled patch against real upstream artifacts.""" + + @unittest.skipUnless( + os.environ.get('MEGATRON_PATCH_BASE_WHEEL'), 'set MEGATRON_PATCH_BASE_WHEEL to the baseline wheel') + def test_baseline_wheel(self): + with tempfile.TemporaryDirectory() as directory: + root = pathlib.Path(directory) + with zipfile.ZipFile(os.environ['MEGATRON_PATCH_BASE_WHEEL']) as wheel: + wheel.extractall(root) + self.verify_complete_patch(root, source=False) + + @unittest.skipUnless( + os.environ.get('MEGATRON_PATCH_BASE_SOURCE'), 'set MEGATRON_PATCH_BASE_SOURCE to baseline sources') + def test_baseline_source(self): + with tempfile.TemporaryDirectory() as directory: + root = pathlib.Path(directory) + base = pathlib.Path(os.environ['MEGATRON_PATCH_BASE_SOURCE']) + for name in installer.patch_files(True): + target = root / name + target.parent.mkdir(parents=True, exist_ok=True) + target.write_bytes((base / name).read_bytes()) + subprocess.run(['git', 'init', '-q', str(root)], check=True) + self.verify_complete_patch(root, source=True) + + def verify_complete_patch(self, root, source): + entries = installer.patch_files(source) + self.assertFalse(installer.is_applied(root)) + installer.apply_patch(root) + installer.apply_patch(root, check_only=True) + installer.apply_patch(root) + for name, diff in entries.items(): + expected = diff.split('index ', 1)[1].split('..', 1)[1].split()[0] + self.assertTrue(blob((root / name).read_text()).startswith(expected), name) + + +if __name__ == '__main__': + unittest.main() diff --git a/tests/test_glm5_hybrid.py b/tests/test_glm5_hybrid.py index 90dab80..85da591 100644 --- a/tests/test_glm5_hybrid.py +++ b/tests/test_glm5_hybrid.py @@ -204,12 +204,13 @@ def test_unpatched_dev_is_rejected_at_glm_boundary(monkeypatch): def test_megatron_patch_is_packaged_and_detectable(): - """The patch must ship inside the package, and carry the marker `is_applied` looks for.""" - from mcore_bridge.tools.apply_megatron_patch import MARKER_FILE, MARKER_SYMBOL, PATCH, is_applied + """The packaged patch must expose runtime targets without requiring Megatron's test tree.""" + from mcore_bridge.tools.apply_megatron_patch import PATCH, is_applied, patch_files - text = PATCH.read_text() - assert f'+++ b/{MARKER_FILE}' in text, f'the patch no longer touches {MARKER_FILE}' - assert MARKER_SYMBOL in text, 'the patch no longer adds the symbol that is_applied() detects' + entries = patch_files(source_checkout=False) + assert 'megatron/core/transformer/transformer_config.py' in entries + assert 'kda_two_stage_gates' in PATCH.read_text() + assert all(name.startswith('megatron/core/') for name in entries) assert not is_applied(PATCH.parent / 'no-such-root')