Files
J-Wash/scripts/fit_worker.py
T
2026-07-13 22:26:50 +02:00

101 lines
3.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import argparse
import json
import logging
import pathlib
import sys
ROOT = pathlib.Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
import config
config.setup_env()
import torch
import transformers
import jlens
class ProgressHandler(logging.Handler):
def emit(self, record):
if not record.args:
return
if record.msg.startswith(" prompt"):
print(
json.dumps(
{
"event": "progress",
"done": record.args[0],
"total": record.args[1],
"seconds": record.args[4],
}
),
flush=True,
)
elif record.msg.startswith(" resuming"):
print(
json.dumps(
{"event": "resume", "done": record.args[0], "total": record.args[1]}
),
flush=True,
)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model", required=True)
parser.add_argument("--device", required=True)
parser.add_argument("--dtype", default="bf16")
parser.add_argument("--quant", default=None)
parser.add_argument("--prompts", required=True)
parser.add_argument("--checkpoint", required=True)
parser.add_argument("--out", required=True)
parser.add_argument("--dim-batch", type=int, default=8)
parser.add_argument("--max-seq-len", type=int, default=128)
parser.add_argument("--source-layers", default=None)
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, handlers=[ProgressHandler()])
prompts = json.loads(pathlib.Path(args.prompts).read_text(encoding="utf-8"))
torch_dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
kwargs = {"dtype": torch_dtype, "device_map": {"": args.device}}
model_source = args.model
if args.quant == "int8":
kwargs["quantization_config"] = transformers.BitsAndBytesConfig(load_in_8bit=True)
elif args.quant == "nf4":
kwargs["quantization_config"] = transformers.BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch_dtype,
bnb_4bit_use_double_quant=True,
)
print(json.dumps({"event": "loading", "model": args.model, "device": args.device}), flush=True)
hf_model = transformers.AutoModelForCausalLM.from_pretrained(model_source, **kwargs)
tokenizer = transformers.AutoTokenizer.from_pretrained(model_source)
model = jlens.from_hf(hf_model, tokenizer)
source_layers = json.loads(args.source_layers) if args.source_layers else None
# large models: the checkpoint (n_layers × d_model² × 4 B) can weigh hundreds
# of MB — writing it after every prompt would wear the SSD for nothing. We
# space it out to target ~150 MB of average writes per prompt.
n_src = len(source_layers) if source_layers else model.n_layers - 1
ckpt_bytes = n_src * model.d_model**2 * 4
checkpoint_every = max(1, round(ckpt_bytes / 150e6))
lens = jlens.fit(
model,
prompts,
source_layers=source_layers,
dim_batch=args.dim_batch,
max_seq_len=args.max_seq_len,
checkpoint_path=args.checkpoint,
checkpoint_every=checkpoint_every,
)
lens.save(args.out)
print(json.dumps({"event": "done", "out": args.out, "n_prompts": lens.n_prompts}), flush=True)
main()