adapterbridge 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- adapterbridge/__init__.py +32 -0
- adapterbridge/cli.py +143 -0
- adapterbridge/core/__init__.py +16 -0
- adapterbridge/core/dry_run.py +96 -0
- adapterbridge/core/inspector.py +53 -0
- adapterbridge/core/lineage.py +85 -0
- adapterbridge/core/remapper.py +97 -0
- adapterbridge/core/template_library.py +75 -0
- adapterbridge/core/template_linter.py +44 -0
- adapterbridge/models/__init__.py +27 -0
- adapterbridge/models/manifest.py +32 -0
- adapterbridge/models/report.py +169 -0
- adapterbridge/models/target_spec.py +40 -0
- adapterbridge/targets/__init__.py +15 -0
- adapterbridge/targets/ollama.py +67 -0
- adapterbridge/targets/registry.py +44 -0
- adapterbridge/targets/sglang.py +88 -0
- adapterbridge/targets/tensorrt.py +68 -0
- adapterbridge/targets/vllm.py +153 -0
- adapterbridge/utils/__init__.py +19 -0
- adapterbridge/utils/hub.py +57 -0
- adapterbridge/utils/jinja_sandbox.py +64 -0
- adapterbridge/utils/safetensors_io.py +96 -0
- adapterbridge-0.1.0.dist-info/METADATA +376 -0
- adapterbridge-0.1.0.dist-info/RECORD +29 -0
- adapterbridge-0.1.0.dist-info/WHEEL +5 -0
- adapterbridge-0.1.0.dist-info/entry_points.txt +8 -0
- adapterbridge-0.1.0.dist-info/licenses/LICENSE +105 -0
- adapterbridge-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""AdapterBridge: LoRA Checkpoint & Config Compatibility Engine for Enterprise Inference Runtimes."""
|
|
2
|
+
|
|
3
|
+
__version__ = "0.1.0"
|
|
4
|
+
|
|
5
|
+
from adapterbridge.models.manifest import CheckpointManifest, TensorMetadata
|
|
6
|
+
from adapterbridge.models.report import (
|
|
7
|
+
IssueSeverity,
|
|
8
|
+
ValidationIssue,
|
|
9
|
+
ValidationReport,
|
|
10
|
+
RemediationAction,
|
|
11
|
+
RemediationOperation,
|
|
12
|
+
RemediationPlan,
|
|
13
|
+
VerificationResult,
|
|
14
|
+
)
|
|
15
|
+
from adapterbridge.models.target_spec import TargetEngine, BaseTargetProfile
|
|
16
|
+
from adapterbridge.core.inspector import AdapterInspector
|
|
17
|
+
|
|
18
|
+
__all__ = [
|
|
19
|
+
"AdapterInspector",
|
|
20
|
+
"CheckpointManifest",
|
|
21
|
+
"TensorMetadata",
|
|
22
|
+
"IssueSeverity",
|
|
23
|
+
"ValidationIssue",
|
|
24
|
+
"ValidationReport",
|
|
25
|
+
"RemediationAction",
|
|
26
|
+
"RemediationOperation",
|
|
27
|
+
"RemediationPlan",
|
|
28
|
+
"VerificationResult",
|
|
29
|
+
"TargetEngine",
|
|
30
|
+
"BaseTargetProfile",
|
|
31
|
+
]
|
|
32
|
+
|
adapterbridge/cli.py
ADDED
|
@@ -0,0 +1,143 @@
|
|
|
1
|
+
"""CLI entry points using Typer and Rich formatting."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
import sys
|
|
6
|
+
from typing import Optional
|
|
7
|
+
import typer
|
|
8
|
+
from rich.console import Console
|
|
9
|
+
from rich.table import Table
|
|
10
|
+
from rich.panel import Panel
|
|
11
|
+
|
|
12
|
+
from adapterbridge.core.inspector import AdapterInspector
|
|
13
|
+
from adapterbridge.models.report import IssueSeverity
|
|
14
|
+
|
|
15
|
+
app = typer.Typer(
|
|
16
|
+
name="adapterbridge",
|
|
17
|
+
help="AdapterBridge: LoRA Checkpoint & Config Compatibility Engine for Enterprise Inference Runtimes.",
|
|
18
|
+
add_completion=False,
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
console = Console()
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@app.command(name="check")
|
|
25
|
+
def check_cmd(
|
|
26
|
+
path: str = typer.Option(..., "--path", "-p", help="Path to checkpoint directory"),
|
|
27
|
+
target: str = typer.Option(..., "--target", "-t", help="Target serving engine (vllm, sglang, ollama, tensorrt)"),
|
|
28
|
+
format: str = typer.Option("text", "--format", "-f", help="Output format: text, json, sarif, markdown, pr-comment"),
|
|
29
|
+
output: Optional[str] = typer.Option(None, "--output", "-o", help="Optional output file destination"),
|
|
30
|
+
):
|
|
31
|
+
"""Static linter and compatibility validator for LoRA checkpoints."""
|
|
32
|
+
try:
|
|
33
|
+
inspector = AdapterInspector(checkpoint_path=path, target_engine=target)
|
|
34
|
+
report = inspector.run_diagnostics()
|
|
35
|
+
except Exception as e:
|
|
36
|
+
console.print(f"[bold red]Error:[/] {str(e)}")
|
|
37
|
+
raise typer.Exit(code=1)
|
|
38
|
+
|
|
39
|
+
if format == "json":
|
|
40
|
+
formatted_output = report.model_dump_json(indent=2)
|
|
41
|
+
elif format == "sarif":
|
|
42
|
+
formatted_output = json.dumps(report.to_sarif(), indent=2)
|
|
43
|
+
elif format == "pr-comment":
|
|
44
|
+
formatted_output = report.to_pr_comment()
|
|
45
|
+
elif format == "markdown":
|
|
46
|
+
status_str = "PASSED" if report.is_compatible else "FAILED"
|
|
47
|
+
md_lines = [
|
|
48
|
+
f"# AdapterBridge Compatibility Report",
|
|
49
|
+
f"**Checkpoint:** `{report.checkpoint_path}` ",
|
|
50
|
+
f"**Target Engine:** `{report.target_engine}` ",
|
|
51
|
+
f"**Status:** `{status_str}` ",
|
|
52
|
+
"",
|
|
53
|
+
"## Diagnostic Issues",
|
|
54
|
+
]
|
|
55
|
+
for issue in report.issues:
|
|
56
|
+
md_lines.append(f"- **[{issue.severity.value.upper()}]** `[{issue.code}]`: {issue.message}")
|
|
57
|
+
formatted_output = "\n".join(md_lines)
|
|
58
|
+
else:
|
|
59
|
+
# Rich Terminal Output
|
|
60
|
+
table = Table(title=f"AdapterBridge Diagnostic Report ({target.upper()})")
|
|
61
|
+
table.add_column("Severity", style="bold")
|
|
62
|
+
table.add_column("Code", style="cyan")
|
|
63
|
+
table.add_column("Message")
|
|
64
|
+
|
|
65
|
+
for issue in report.issues:
|
|
66
|
+
color = "red" if issue.severity == IssueSeverity.ERROR else "yellow"
|
|
67
|
+
table.add_row(f"[{color}]{issue.severity.value.upper()}[/{color}]", issue.code, issue.message)
|
|
68
|
+
|
|
69
|
+
console.print(table)
|
|
70
|
+
status_panel_style = "green" if report.is_compatible else "red"
|
|
71
|
+
console.print(Panel(report.summary, style=status_panel_style))
|
|
72
|
+
formatted_output = None
|
|
73
|
+
|
|
74
|
+
if output and formatted_output:
|
|
75
|
+
with open(output, "w", encoding="utf-8") as f:
|
|
76
|
+
f.write(formatted_output)
|
|
77
|
+
console.print(f"[bold green]Report written to:[/] {output}")
|
|
78
|
+
elif formatted_output and format != "text":
|
|
79
|
+
console.print(formatted_output)
|
|
80
|
+
|
|
81
|
+
if not report.is_compatible:
|
|
82
|
+
raise typer.Exit(code=1)
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
@app.command(name="fix")
|
|
86
|
+
def fix_cmd(
|
|
87
|
+
src: str = typer.Option(..., "--src", "-s", help="Source checkpoint path"),
|
|
88
|
+
dst: str = typer.Option(..., "--dst", "-d", help="Destination path for repaired checkpoint"),
|
|
89
|
+
target: str = typer.Option(..., "--target", "-t", help="Target engine (vllm, sglang, ollama, tensorrt)"),
|
|
90
|
+
base_model: Optional[str] = typer.Option(None, "--base-model", "-b", help="Fallback base model ID"),
|
|
91
|
+
):
|
|
92
|
+
"""Automated normalizer & metadata synthesizer for LoRA checkpoints."""
|
|
93
|
+
try:
|
|
94
|
+
inspector = AdapterInspector(checkpoint_path=src, target_engine=target)
|
|
95
|
+
console.print(f"[bold blue]Analyzing source checkpoint:[/] {src}")
|
|
96
|
+
plan = inspector.auto_repair(destination_path=dst, fallback_base_model=base_model)
|
|
97
|
+
|
|
98
|
+
console.print(f"[bold green]Repaired checkpoint saved to:[/] {plan.output_path}")
|
|
99
|
+
console.print(f"Executed {len(plan.operations)} remediation operation(s).")
|
|
100
|
+
except Exception as e:
|
|
101
|
+
console.print(f"[bold red]Remediation Error:[/] {str(e)}")
|
|
102
|
+
raise typer.Exit(code=1)
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
@app.command(name="verify")
|
|
106
|
+
def verify_cmd(
|
|
107
|
+
path: str = typer.Option(..., "--path", "-p", help="Path to checkpoint directory"),
|
|
108
|
+
target: str = typer.Option(..., "--target", "-t", help="Target engine"),
|
|
109
|
+
tensor_parallel_size: int = typer.Option(1, "--tensor-parallel-size", "-tp", help="Tensor Parallelism degree"),
|
|
110
|
+
):
|
|
111
|
+
"""Zero-GPU mock serving dry-run engine."""
|
|
112
|
+
try:
|
|
113
|
+
inspector = AdapterInspector(checkpoint_path=path, target_engine=target)
|
|
114
|
+
res = inspector.verify_dry_run(tensor_parallel_size=tensor_parallel_size)
|
|
115
|
+
|
|
116
|
+
if res.success:
|
|
117
|
+
console.print(Panel(f"SUCCESS: {res.reason}\nEstimated Adapter RAM: {res.memory_estimate_mb} MB", style="green"))
|
|
118
|
+
else:
|
|
119
|
+
console.print(Panel(f"FAILURE: {res.reason}\nIssues: {res.sharding_issues}", style="red"))
|
|
120
|
+
raise typer.Exit(code=1)
|
|
121
|
+
except Exception as e:
|
|
122
|
+
console.print(f"[bold red]Verification Error:[/] {str(e)}")
|
|
123
|
+
raise typer.Exit(code=1)
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
@app.command(name="export")
|
|
127
|
+
def export_cmd(
|
|
128
|
+
path: str = typer.Option(..., "--path", "-p", help="Path to checkpoint directory"),
|
|
129
|
+
target: str = typer.Option(..., "--target", "-t", help="Target engine"),
|
|
130
|
+
output: str = typer.Option(..., "--output", "-o", help="Export target destination directory"),
|
|
131
|
+
):
|
|
132
|
+
"""Serving-targeted bundler."""
|
|
133
|
+
try:
|
|
134
|
+
inspector = AdapterInspector(checkpoint_path=path, target_engine=target)
|
|
135
|
+
plan = inspector.auto_repair(destination_path=output)
|
|
136
|
+
console.print(f"[bold green]Successfully exported checkpoint bundle to:[/] {output}")
|
|
137
|
+
except Exception as e:
|
|
138
|
+
console.print(f"[bold red]Export Error:[/] {str(e)}")
|
|
139
|
+
raise typer.Exit(code=1)
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
if __name__ == "__main__":
|
|
143
|
+
app()
|
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
"""Core subsystems package exports."""
|
|
2
|
+
|
|
3
|
+
from adapterbridge.core.inspector import AdapterInspector
|
|
4
|
+
from adapterbridge.core.lineage import scan_checkpoint
|
|
5
|
+
from adapterbridge.core.remapper import execute_remediation_plan, normalize_tensor_key
|
|
6
|
+
from adapterbridge.core.template_linter import lint_chat_template
|
|
7
|
+
from adapterbridge.core.dry_run import run_dry_run_verification
|
|
8
|
+
|
|
9
|
+
__all__ = [
|
|
10
|
+
"AdapterInspector",
|
|
11
|
+
"scan_checkpoint",
|
|
12
|
+
"execute_remediation_plan",
|
|
13
|
+
"normalize_tensor_key",
|
|
14
|
+
"lint_chat_template",
|
|
15
|
+
"run_dry_run_verification",
|
|
16
|
+
]
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
"""Zero-GPU mock serving dry-run engine supporting PyTorch meta-tensors and NumPy shape math."""
|
|
2
|
+
|
|
3
|
+
from typing import Dict, List, Optional
|
|
4
|
+
from adapterbridge.models.manifest import CheckpointManifest
|
|
5
|
+
from adapterbridge.models.report import VerificationResult
|
|
6
|
+
|
|
7
|
+
# Optional PyTorch import check
|
|
8
|
+
try:
|
|
9
|
+
import torch
|
|
10
|
+
HAS_TORCH = True
|
|
11
|
+
except ImportError:
|
|
12
|
+
HAS_TORCH = False
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
DTYPE_BYTES: Dict[str, float] = {
|
|
16
|
+
"float32": 4,
|
|
17
|
+
"fp32": 4,
|
|
18
|
+
"float16": 2,
|
|
19
|
+
"fp16": 2,
|
|
20
|
+
"bfloat16": 2,
|
|
21
|
+
"bf16": 2,
|
|
22
|
+
"int8": 1,
|
|
23
|
+
"int4": 0.5,
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def run_dry_run_verification(
|
|
28
|
+
manifest: CheckpointManifest, tensor_parallel_size: int = 1
|
|
29
|
+
) -> VerificationResult:
|
|
30
|
+
"""Simulate target-runtime memory allocation, rank scaling, and tensor-parallel sharding math."""
|
|
31
|
+
sharding_issues: List[str] = []
|
|
32
|
+
total_bytes = 0
|
|
33
|
+
used_meta_tensors = False
|
|
34
|
+
|
|
35
|
+
if HAS_TORCH:
|
|
36
|
+
try:
|
|
37
|
+
used_meta_tensors = True
|
|
38
|
+
for t_name, t_meta in manifest.tensor_manifest.items():
|
|
39
|
+
# Materialize meta tensor without allocating GPU or host RAM
|
|
40
|
+
meta_t = torch.empty(t_meta.shape, device="meta")
|
|
41
|
+
element_bytes = meta_t.element_size() if meta_t.element_size() > 0 else 2
|
|
42
|
+
total_bytes += int(t_meta.n_elements * element_bytes)
|
|
43
|
+
|
|
44
|
+
# Validate Tensor Parallelism sharding division
|
|
45
|
+
if tensor_parallel_size > 1 and meta_t.dim() >= 2:
|
|
46
|
+
major_dim = max(meta_t.shape)
|
|
47
|
+
if major_dim % tensor_parallel_size != 0:
|
|
48
|
+
sharding_issues.append(
|
|
49
|
+
f"Tensor '{t_name}' dimension {major_dim} cannot be split evenly across tensor-parallel-size={tensor_parallel_size}."
|
|
50
|
+
)
|
|
51
|
+
except Exception:
|
|
52
|
+
used_meta_tensors = False
|
|
53
|
+
|
|
54
|
+
if not used_meta_tensors:
|
|
55
|
+
# Fallback NumPy shape math
|
|
56
|
+
for t_name, t_meta in manifest.tensor_manifest.items():
|
|
57
|
+
element_bytes = DTYPE_BYTES.get(t_meta.dtype.lower(), 2)
|
|
58
|
+
total_bytes += int(t_meta.n_elements * element_bytes)
|
|
59
|
+
|
|
60
|
+
if tensor_parallel_size > 1 and len(t_meta.shape) >= 2:
|
|
61
|
+
major_dim = max(t_meta.shape)
|
|
62
|
+
if major_dim % tensor_parallel_size != 0:
|
|
63
|
+
sharding_issues.append(
|
|
64
|
+
f"Tensor '{t_name}' dimension {major_dim} cannot be split evenly across tensor-parallel-size={tensor_parallel_size}."
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
# Validate LoRA rank dimension matching
|
|
68
|
+
for t_name, t_meta in manifest.tensor_manifest.items():
|
|
69
|
+
if ("lora_A" in t_name or "lora_B" in t_name) and manifest.lora_r:
|
|
70
|
+
if manifest.lora_r not in t_meta.shape:
|
|
71
|
+
sharding_issues.append(
|
|
72
|
+
f"Tensor '{t_name}' shape {t_meta.shape} does not contain configured rank r={manifest.lora_r}."
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
memory_estimate_mb = round(total_bytes / (1024 * 1024), 3)
|
|
76
|
+
|
|
77
|
+
scaling_info = ""
|
|
78
|
+
if manifest.lora_r and manifest.lora_alpha:
|
|
79
|
+
scale = manifest.lora_alpha / manifest.lora_r
|
|
80
|
+
scaling_info = f" Rank r={manifest.lora_r}, alpha={manifest.lora_alpha}, scale={scale:.2f}."
|
|
81
|
+
|
|
82
|
+
mode_info = " [PyTorch meta-tensors]" if used_meta_tensors else " [NumPy shape engine]"
|
|
83
|
+
success = len(sharding_issues) == 0
|
|
84
|
+
reason = (
|
|
85
|
+
f"Dry-run simulation passed successfully.{scaling_info}{mode_info}"
|
|
86
|
+
if success
|
|
87
|
+
else f"Dry-run failed with {len(sharding_issues)} sharding/shape issue(s).{mode_info}"
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
return VerificationResult(
|
|
91
|
+
success=success,
|
|
92
|
+
reason=reason,
|
|
93
|
+
tensor_parallel_size=tensor_parallel_size,
|
|
94
|
+
memory_estimate_mb=memory_estimate_mb,
|
|
95
|
+
sharding_issues=sharding_issues,
|
|
96
|
+
)
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
"""Top-level SDK orchestration for checkpoint inspection, repair, and dry-run verification."""
|
|
2
|
+
|
|
3
|
+
from typing import Optional, Union
|
|
4
|
+
from adapterbridge.core.dry_run import run_dry_run_verification
|
|
5
|
+
from adapterbridge.core.lineage import scan_checkpoint
|
|
6
|
+
from adapterbridge.core.remapper import execute_remediation_plan
|
|
7
|
+
from adapterbridge.models.manifest import CheckpointManifest
|
|
8
|
+
from adapterbridge.models.report import RemediationPlan, ValidationReport, VerificationResult
|
|
9
|
+
from adapterbridge.models.target_spec import TargetEngine
|
|
10
|
+
from adapterbridge.targets.registry import get_target_profile
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class AdapterInspector:
|
|
14
|
+
"""Primary SDK entry point for evaluating and remediating LoRA checkpoints."""
|
|
15
|
+
|
|
16
|
+
def __init__(self, checkpoint_path: str, target_engine: Union[TargetEngine, str]):
|
|
17
|
+
self.checkpoint_path = checkpoint_path
|
|
18
|
+
self.target_engine_str = target_engine.value if isinstance(target_engine, TargetEngine) else str(target_engine)
|
|
19
|
+
self.target_profile = get_target_profile(self.target_engine_str)
|
|
20
|
+
self._manifest: Optional[CheckpointManifest] = None
|
|
21
|
+
|
|
22
|
+
@property
|
|
23
|
+
def manifest(self) -> CheckpointManifest:
|
|
24
|
+
"""Lazy-loaded CheckpointManifest."""
|
|
25
|
+
if self._manifest is None:
|
|
26
|
+
self._manifest = scan_checkpoint(self.checkpoint_path)
|
|
27
|
+
return self._manifest
|
|
28
|
+
|
|
29
|
+
def run_diagnostics(self) -> ValidationReport:
|
|
30
|
+
"""Perform static schema inspection and validation against the target profile."""
|
|
31
|
+
return self.target_profile.validate_manifest(self.manifest)
|
|
32
|
+
|
|
33
|
+
def auto_repair(
|
|
34
|
+
self, destination_path: str, fallback_base_model: Optional[str] = None
|
|
35
|
+
) -> RemediationPlan:
|
|
36
|
+
"""Generate and execute a non-destructive remediation plan."""
|
|
37
|
+
current_manifest = self.manifest
|
|
38
|
+
if fallback_base_model and not current_manifest.base_model_id:
|
|
39
|
+
current_manifest = current_manifest.model_copy(
|
|
40
|
+
update={"base_model_id": fallback_base_model}
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
plan = self.target_profile.generate_remediation_plan(
|
|
44
|
+
manifest=current_manifest, destination_path=destination_path
|
|
45
|
+
)
|
|
46
|
+
if plan.is_executable and plan.operations:
|
|
47
|
+
execute_remediation_plan(current_manifest, plan)
|
|
48
|
+
|
|
49
|
+
return plan
|
|
50
|
+
|
|
51
|
+
def verify_dry_run(self, tensor_parallel_size: int = 1) -> VerificationResult:
|
|
52
|
+
"""Perform zero-GPU mock dry-run simulation."""
|
|
53
|
+
return run_dry_run_verification(self.manifest, tensor_parallel_size=tensor_parallel_size)
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
"""Checkpoint inspection scanner and metadata lineage resolution."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
from typing import Dict
|
|
6
|
+
from adapterbridge.models.manifest import CheckpointManifest, TensorMetadata
|
|
7
|
+
from adapterbridge.utils.safetensors_io import extract_tensor_headers
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def scan_checkpoint(checkpoint_path: str) -> CheckpointManifest:
|
|
11
|
+
"""Inspect a local checkpoint directory and construct an immutable CheckpointManifest."""
|
|
12
|
+
if not os.path.exists(checkpoint_path):
|
|
13
|
+
return CheckpointManifest(
|
|
14
|
+
checkpoint_path=checkpoint_path,
|
|
15
|
+
is_adapter=False,
|
|
16
|
+
has_config=False,
|
|
17
|
+
has_tokenizer=False,
|
|
18
|
+
has_chat_template=False,
|
|
19
|
+
tensor_manifest={},
|
|
20
|
+
validation_errors=[f"Checkpoint path '{checkpoint_path}' does not exist."],
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
adapter_config_path = os.path.join(checkpoint_path, "adapter_config.json")
|
|
24
|
+
config_path = os.path.join(checkpoint_path, "config.json")
|
|
25
|
+
tok_config_path = os.path.join(checkpoint_path, "tokenizer_config.json")
|
|
26
|
+
tok_json_path = os.path.join(checkpoint_path, "tokenizer.json")
|
|
27
|
+
|
|
28
|
+
is_adapter = os.path.exists(adapter_config_path)
|
|
29
|
+
has_config = os.path.exists(config_path)
|
|
30
|
+
has_tokenizer = os.path.exists(tok_config_path) or os.path.exists(tok_json_path)
|
|
31
|
+
|
|
32
|
+
base_model_id = None
|
|
33
|
+
adapter_type = "lora"
|
|
34
|
+
lora_r = None
|
|
35
|
+
lora_alpha = None
|
|
36
|
+
target_modules = []
|
|
37
|
+
|
|
38
|
+
if is_adapter:
|
|
39
|
+
try:
|
|
40
|
+
with open(adapter_config_path, "r", encoding="utf-8") as f:
|
|
41
|
+
ac_data = json.load(f)
|
|
42
|
+
base_model_id = ac_data.get("base_model_name_or_path")
|
|
43
|
+
adapter_type = ac_data.get("peft_type", "lora").lower()
|
|
44
|
+
lora_r = ac_data.get("r")
|
|
45
|
+
lora_alpha = ac_data.get("lora_alpha")
|
|
46
|
+
t_mods = ac_data.get("target_modules", [])
|
|
47
|
+
if isinstance(t_mods, list):
|
|
48
|
+
target_modules = t_mods
|
|
49
|
+
elif isinstance(t_mods, str):
|
|
50
|
+
target_modules = [t_mods]
|
|
51
|
+
except Exception:
|
|
52
|
+
pass
|
|
53
|
+
|
|
54
|
+
has_chat_template = False
|
|
55
|
+
if os.path.exists(tok_config_path):
|
|
56
|
+
try:
|
|
57
|
+
with open(tok_config_path, "r", encoding="utf-8") as f:
|
|
58
|
+
tc_data = json.load(f)
|
|
59
|
+
if "chat_template" in tc_data and tc_data["chat_template"]:
|
|
60
|
+
has_chat_template = True
|
|
61
|
+
except Exception:
|
|
62
|
+
pass
|
|
63
|
+
|
|
64
|
+
tensor_manifest: Dict[str, TensorMetadata] = {}
|
|
65
|
+
for root, _, files in os.walk(checkpoint_path):
|
|
66
|
+
for file in files:
|
|
67
|
+
if file.endswith(".safetensors"):
|
|
68
|
+
full_path = os.path.join(root, file)
|
|
69
|
+
headers = extract_tensor_headers(full_path)
|
|
70
|
+
tensor_manifest.update(headers)
|
|
71
|
+
|
|
72
|
+
return CheckpointManifest(
|
|
73
|
+
checkpoint_path=checkpoint_path,
|
|
74
|
+
is_adapter=is_adapter,
|
|
75
|
+
base_model_id=base_model_id,
|
|
76
|
+
adapter_type=adapter_type,
|
|
77
|
+
lora_r=lora_r,
|
|
78
|
+
lora_alpha=lora_alpha,
|
|
79
|
+
target_modules=target_modules,
|
|
80
|
+
has_config=has_config,
|
|
81
|
+
has_tokenizer=has_tokenizer,
|
|
82
|
+
has_chat_template=has_chat_template,
|
|
83
|
+
tensor_manifest=tensor_manifest,
|
|
84
|
+
validation_errors=[],
|
|
85
|
+
)
|
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
"""Tensor state-dict key normalization and remediation plan execution engine."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
import shutil
|
|
6
|
+
import uuid
|
|
7
|
+
from typing import Dict
|
|
8
|
+
from adapterbridge.models.manifest import CheckpointManifest
|
|
9
|
+
from adapterbridge.models.report import RemediationAction, RemediationPlan
|
|
10
|
+
from adapterbridge.utils.safetensors_io import remap_safetensors_file
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
KNOWN_PREFIX_REMAPS = [
|
|
14
|
+
("base_model.model.model.layers.", "model.layers."),
|
|
15
|
+
("base_model.model.", "model."),
|
|
16
|
+
]
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def normalize_tensor_key(key: str) -> str:
|
|
20
|
+
"""Normalize a tensor key using standard framework prefix stripping."""
|
|
21
|
+
for old_prefix, new_prefix in KNOWN_PREFIX_REMAPS:
|
|
22
|
+
if key.startswith(old_prefix):
|
|
23
|
+
return key.replace(old_prefix, new_prefix, 1)
|
|
24
|
+
return key
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def execute_remediation_plan(manifest: CheckpointManifest, plan: RemediationPlan) -> str:
|
|
28
|
+
"""Execute a remediation plan using atomic staging directories.
|
|
29
|
+
|
|
30
|
+
Returns the final destination directory path.
|
|
31
|
+
"""
|
|
32
|
+
dst_path = os.path.abspath(plan.output_path)
|
|
33
|
+
staging_dir = os.path.join(
|
|
34
|
+
os.path.dirname(dst_path),
|
|
35
|
+
f".adapterbridge_staging_{uuid.uuid4().hex[:8]}"
|
|
36
|
+
)
|
|
37
|
+
os.makedirs(staging_dir, exist_ok=True)
|
|
38
|
+
|
|
39
|
+
try:
|
|
40
|
+
# Copy source files to staging
|
|
41
|
+
if os.path.exists(manifest.checkpoint_path):
|
|
42
|
+
for item in os.listdir(manifest.checkpoint_path):
|
|
43
|
+
s_item = os.path.join(manifest.checkpoint_path, item)
|
|
44
|
+
d_item = os.path.join(staging_dir, item)
|
|
45
|
+
if os.path.isfile(s_item):
|
|
46
|
+
shutil.copy2(s_item, d_item)
|
|
47
|
+
elif os.path.isdir(s_item):
|
|
48
|
+
shutil.copytree(s_item, d_item)
|
|
49
|
+
|
|
50
|
+
# Execute operations
|
|
51
|
+
for op in plan.operations:
|
|
52
|
+
if op.action == RemediationAction.SYNTHESIZE_FILE:
|
|
53
|
+
target_rel = os.path.relpath(op.target_path, plan.output_path)
|
|
54
|
+
staged_file = os.path.join(staging_dir, target_rel)
|
|
55
|
+
os.makedirs(os.path.dirname(staged_file), exist_ok=True)
|
|
56
|
+
|
|
57
|
+
if staged_file.endswith(".json"):
|
|
58
|
+
with open(staged_file, "w", encoding="utf-8") as f:
|
|
59
|
+
json.dump(op.details, f, indent=2)
|
|
60
|
+
elif staged_file.endswith("Modelfile"):
|
|
61
|
+
with open(staged_file, "w", encoding="utf-8") as f:
|
|
62
|
+
f.write(f"FROM {op.details.get('base_model', 'llama3.1')}\n")
|
|
63
|
+
f.write(f"ADAPTER {op.details.get('adapter_path', './adapter_model.safetensors')}\n")
|
|
64
|
+
|
|
65
|
+
elif op.action == RemediationAction.INJECT_CHAT_TEMPLATE:
|
|
66
|
+
tok_config_path = os.path.join(staging_dir, "tokenizer_config.json")
|
|
67
|
+
tc_data: Dict[str, str] = {}
|
|
68
|
+
if os.path.exists(tok_config_path):
|
|
69
|
+
try:
|
|
70
|
+
with open(tok_config_path, "r", encoding="utf-8") as f:
|
|
71
|
+
tc_data = json.load(f)
|
|
72
|
+
except Exception:
|
|
73
|
+
pass
|
|
74
|
+
tc_data["chat_template"] = op.details.get("chat_template", "")
|
|
75
|
+
with open(tok_config_path, "w", encoding="utf-8") as f:
|
|
76
|
+
json.dump(tc_data, f, indent=2)
|
|
77
|
+
|
|
78
|
+
elif op.action == RemediationAction.REMAP_TENSOR_KEY:
|
|
79
|
+
key_map = op.details.get("key_map", {})
|
|
80
|
+
for root, _, files in os.walk(staging_dir):
|
|
81
|
+
for file in files:
|
|
82
|
+
if file.endswith(".safetensors"):
|
|
83
|
+
staged_st = os.path.join(root, file)
|
|
84
|
+
temp_st = staged_st + ".tmp"
|
|
85
|
+
remap_safetensors_file(staged_st, temp_st, key_map)
|
|
86
|
+
os.replace(temp_st, staged_st)
|
|
87
|
+
|
|
88
|
+
# Atomic rename into final destination
|
|
89
|
+
if os.path.exists(dst_path):
|
|
90
|
+
shutil.rmtree(dst_path)
|
|
91
|
+
shutil.move(staging_dir, dst_path)
|
|
92
|
+
return dst_path
|
|
93
|
+
|
|
94
|
+
except Exception as e:
|
|
95
|
+
if os.path.exists(staging_dir):
|
|
96
|
+
shutil.rmtree(staging_dir, ignore_errors=True)
|
|
97
|
+
raise RuntimeError(f"Remediation execution failed: {str(e)}") from e
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
"""Canonical Jinja2 chat template library for supported model architectures."""
|
|
2
|
+
|
|
3
|
+
from typing import Dict, Optional
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
CANONICAL_TEMPLATES: Dict[str, str] = {
|
|
7
|
+
# Llama 3 / 3.1 / 3.3
|
|
8
|
+
"llama": (
|
|
9
|
+
"{% for message in messages %}"
|
|
10
|
+
"{{ '<|start_header_id|>' + message['role'] + '<|end_header_id|>\n\n' + message['content'] + '<|eot_id|>' }}"
|
|
11
|
+
"{% endfor %}"
|
|
12
|
+
"{% if add_generation_prompt %}"
|
|
13
|
+
"{{ '<|start_header_id|>assistant<|end_header_id|>\n\n' }}"
|
|
14
|
+
"{% endif %}"
|
|
15
|
+
),
|
|
16
|
+
# Qwen 2 / 2.5 / ChatML
|
|
17
|
+
"qwen2": (
|
|
18
|
+
"{% for message in messages %}"
|
|
19
|
+
"{{ '<|im_start|>' + message['role'] + '\n' + message['content'] + '<|im_end|>\n' }}"
|
|
20
|
+
"{% endfor %}"
|
|
21
|
+
"{% if add_generation_prompt %}"
|
|
22
|
+
"{{ '<|im_start|>assistant\n' }}"
|
|
23
|
+
"{% endif %}"
|
|
24
|
+
),
|
|
25
|
+
# Mistral / Mixtral
|
|
26
|
+
"mistral": (
|
|
27
|
+
"{% for message in messages %}"
|
|
28
|
+
"{% if message['role'] == 'user' %}"
|
|
29
|
+
"{{ '[INST] ' + message['content'] + ' [/INST]' }}"
|
|
30
|
+
"{% else %}"
|
|
31
|
+
"{{ message['content'] }}"
|
|
32
|
+
"{% endif %}"
|
|
33
|
+
"{% endfor %}"
|
|
34
|
+
),
|
|
35
|
+
# DeepSeek R1 / V3
|
|
36
|
+
"deepseek": (
|
|
37
|
+
"{% for message in messages %}"
|
|
38
|
+
"{% if message['role'] == 'user' %}"
|
|
39
|
+
"{{ '<|User|>' + message['content'] }}"
|
|
40
|
+
"{% elif message['role'] == 'assistant' %}"
|
|
41
|
+
"{{ '<|Assistant|>' + message['content'] }}"
|
|
42
|
+
"{% endif %}"
|
|
43
|
+
"{% endfor %}"
|
|
44
|
+
"{% if add_generation_prompt %}"
|
|
45
|
+
"{{ '<|Assistant|>' }}"
|
|
46
|
+
"{% endif %}"
|
|
47
|
+
),
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
# Aliases for architectural model_type strings
|
|
51
|
+
ARCHITECTURE_ALIASES: Dict[str, str] = {
|
|
52
|
+
"llama": "llama",
|
|
53
|
+
"llamaforcausalm": "llama",
|
|
54
|
+
"qwen2": "qwen2",
|
|
55
|
+
"qwen2forcausallm": "qwen2",
|
|
56
|
+
"mistral": "mistral",
|
|
57
|
+
"mistralforcausallm": "mistral",
|
|
58
|
+
"mixtral": "mistral",
|
|
59
|
+
"deepseek": "deepseek",
|
|
60
|
+
"deepseekv2": "deepseek",
|
|
61
|
+
"deepseekv3": "deepseek",
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def get_canonical_template_for_architecture(architecture_or_model_type: Optional[str]) -> str:
|
|
66
|
+
"""Return canonical Jinja2 chat template for given architecture or model_type.
|
|
67
|
+
|
|
68
|
+
Defaults to Llama-3 format if model type is missing or unknown.
|
|
69
|
+
"""
|
|
70
|
+
if not architecture_or_model_type:
|
|
71
|
+
return CANONICAL_TEMPLATES["llama"]
|
|
72
|
+
|
|
73
|
+
normalized = str(architecture_or_model_type).lower().replace("_", "").replace("-", "")
|
|
74
|
+
family = ARCHITECTURE_ALIASES.get(normalized, "llama")
|
|
75
|
+
return CANONICAL_TEMPLATES.get(family, CANONICAL_TEMPLATES["llama"])
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
"""Chat template linter and evaluation test vectors."""
|
|
2
|
+
|
|
3
|
+
from typing import List, Tuple
|
|
4
|
+
from adapterbridge.utils.jinja_sandbox import render_template_sandboxed, validate_template_syntax
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
TEST_CONVERSATION_VECTORS = [
|
|
8
|
+
# 1. Simple user prompt
|
|
9
|
+
[{"role": "user", "content": "Hello, world!"}],
|
|
10
|
+
# 2. Multi-turn system + user + assistant
|
|
11
|
+
[
|
|
12
|
+
{"role": "system", "content": "You are a helpful assistant."},
|
|
13
|
+
{"role": "user", "content": "What is Python?"},
|
|
14
|
+
{"role": "assistant", "content": "Python is a programming language."},
|
|
15
|
+
{"role": "user", "content": "Tell me more."},
|
|
16
|
+
],
|
|
17
|
+
]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def lint_chat_template(template_str: str) -> Tuple[bool, List[str]]:
|
|
21
|
+
"""Validate a Jinja2 chat template against standard conversation test vectors.
|
|
22
|
+
|
|
23
|
+
Returns (is_valid, list_of_errors).
|
|
24
|
+
"""
|
|
25
|
+
valid_syntax, err = validate_template_syntax(template_str)
|
|
26
|
+
if not valid_syntax:
|
|
27
|
+
return False, [err]
|
|
28
|
+
|
|
29
|
+
errors: List[str] = []
|
|
30
|
+
for idx, vector in enumerate(TEST_CONVERSATION_VECTORS, start=1):
|
|
31
|
+
success, res = render_template_sandboxed(
|
|
32
|
+
template_str=template_str,
|
|
33
|
+
messages=vector,
|
|
34
|
+
add_generation_prompt=True,
|
|
35
|
+
timeout_seconds=2.0,
|
|
36
|
+
)
|
|
37
|
+
if not success:
|
|
38
|
+
errors.append(f"Vector #{idx} render error: {res}")
|
|
39
|
+
break
|
|
40
|
+
elif not res or len(res.strip()) == 0:
|
|
41
|
+
errors.append(f"Vector #{idx} produced empty output.")
|
|
42
|
+
break
|
|
43
|
+
|
|
44
|
+
return len(errors) == 0, errors
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
"""Model package exports."""
|
|
2
|
+
|
|
3
|
+
from adapterbridge.models.manifest import CheckpointManifest, TensorMetadata
|
|
4
|
+
from adapterbridge.models.report import (
|
|
5
|
+
IssueSeverity,
|
|
6
|
+
ValidationIssue,
|
|
7
|
+
ValidationReport,
|
|
8
|
+
RemediationAction,
|
|
9
|
+
RemediationOperation,
|
|
10
|
+
RemediationPlan,
|
|
11
|
+
VerificationResult,
|
|
12
|
+
)
|
|
13
|
+
from adapterbridge.models.target_spec import TargetEngine, BaseTargetProfile
|
|
14
|
+
|
|
15
|
+
__all__ = [
|
|
16
|
+
"CheckpointManifest",
|
|
17
|
+
"TensorMetadata",
|
|
18
|
+
"IssueSeverity",
|
|
19
|
+
"ValidationIssue",
|
|
20
|
+
"ValidationReport",
|
|
21
|
+
"RemediationAction",
|
|
22
|
+
"RemediationOperation",
|
|
23
|
+
"RemediationPlan",
|
|
24
|
+
"VerificationResult",
|
|
25
|
+
"TargetEngine",
|
|
26
|
+
"BaseTargetProfile",
|
|
27
|
+
]
|