lucataco commited on
Commit
d004817
·
verified ·
1 Parent(s): 953e003

clef_mlx.py: use the repo name (not the cache snapshot hash) as the default model name

Browse files
Files changed (1) hide show
  1. clef_mlx.py +10 -2
clef_mlx.py CHANGED
@@ -590,6 +590,14 @@ def _default_model() -> str | None:
590
  return str(here) if (here / "joint_head.safetensors").exists() else None
591
 
592
 
 
 
 
 
 
 
 
 
593
  def _json_arg(value: str) -> Any:
594
  """A JSON string, a path to a JSON file, or '-' for stdin."""
595
  if value == "-":
@@ -646,7 +654,7 @@ def _cmd_predict(args) -> None:
646
  if args.state is None or args.questions is None:
647
  sys.exit("predict needs --state and --questions, or --request")
648
  request = {"state": _state_arg(args.state), "questions": _json_arg(args.questions)}
649
- request.setdefault("model", args.name or Path(args.model).name)
650
  if args.image:
651
  from PIL import Image
652
 
@@ -665,7 +673,7 @@ def _cmd_serve(args) -> None:
665
  from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
666
 
667
  model = load(args.model)
668
- served_name = args.name or Path(args.model).name
669
  model.systemone({"model": served_name, "state": "warmup",
670
  "questions": {"w": {"type": "noul", "instructions": "Is this a warmup?"}}})
671
  lock = threading.Lock() # one GPU: run requests one at a time
 
590
  return str(here) if (here / "joint_head.safetensors").exists() else None
591
 
592
 
593
+ def _model_name(model: str) -> str:
594
+ """Short name for responses: 'clef-flash-4bit' for a repo id, a local dir, or a Hub cache snapshot."""
595
+ path = Path(model)
596
+ if path.parent.name == "snapshots" and path.parent.parent.name.startswith("models--"):
597
+ return path.parent.parent.name.split("--")[-1]
598
+ return path.name
599
+
600
+
601
  def _json_arg(value: str) -> Any:
602
  """A JSON string, a path to a JSON file, or '-' for stdin."""
603
  if value == "-":
 
654
  if args.state is None or args.questions is None:
655
  sys.exit("predict needs --state and --questions, or --request")
656
  request = {"state": _state_arg(args.state), "questions": _json_arg(args.questions)}
657
+ request.setdefault("model", args.name or _model_name(args.model))
658
  if args.image:
659
  from PIL import Image
660
 
 
673
  from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
674
 
675
  model = load(args.model)
676
+ served_name = args.name or _model_name(args.model)
677
  model.systemone({"model": served_name, "state": "warmup",
678
  "questions": {"w": {"type": "noul", "instructions": "Is this a warmup?"}}})
679
  lock = threading.Lock() # one GPU: run requests one at a time