Parse command line arguments
()
| 41 | |
| 42 | |
| 43 | def parse_args(): |
| 44 | """Parse command line arguments""" |
| 45 | parser = argparse.ArgumentParser(description="Extract LoRA weights from the difference between source and target models", formatter_class=argparse.ArgumentDefaultsHelpFormatter) |
| 46 | |
| 47 | # Source model parameters |
| 48 | parser.add_argument("--source-model", type=str, required=True, help="Path to source model") |
| 49 | parser.add_argument("--source-type", type=str, choices=["safetensors", "pytorch"], default="safetensors", help="Source model format type") |
| 50 | |
| 51 | # Target model parameters |
| 52 | parser.add_argument("--target-model", type=str, required=True, help="Path to target model (fine-tuned model)") |
| 53 | parser.add_argument("--target-type", type=str, choices=["safetensors", "pytorch"], default="safetensors", help="Target model format type") |
| 54 | |
| 55 | # Output parameters |
| 56 | parser.add_argument("--output", type=str, required=True, help="Path to output LoRA model") |
| 57 | parser.add_argument("--output-format", type=str, choices=["safetensors", "pytorch"], default="safetensors", help="Output LoRA model format") |
| 58 | |
| 59 | # LoRA related parameters |
| 60 | parser.add_argument("--rank", type=int, default=32, help="LoRA rank value") |
| 61 | |
| 62 | parser.add_argument("--output-dtype", type=str, choices=["float32", "fp32", "float16", "fp16", "bfloat16", "bf16"], default="bf16", help="Output weight data type") |
| 63 | parser.add_argument("--diff-only", action="store_true", help="Save all weights as direct diff without LoRA decomposition") |
| 64 | |
| 65 | return parser.parse_args() |
| 66 | |
| 67 | |
| 68 | def load_model_weights(model_path: str, model_type: str) -> Dict[str, torch.Tensor]: |