hymenjj/llama-cpp-python-prebuilt
0
1from __future__ import annotations2 3import argparse4 5from typing import List, Literal, Union, Any, Type, TypeVar6 7from pydantic import BaseModel8 9 10def _get_base_type(annotation: Type[Any]) -> Type[Any]:11 if getattr(annotation, "__origin__", None) is Literal:12 assert hasattr(annotation, "__args__") and len(annotation.__args__) >= 1 # type: ignore13 return type(annotation.__args__[0]) # type: ignore14 elif getattr(annotation, "__origin__", None) is Union:15 assert hasattr(annotation, "__args__") and len(annotation.__args__) >= 1 # type: ignore16 non_optional_args: List[Type[Any]] = [17 arg for arg in annotation.__args__ if arg is not type(None) # type: ignore18 ]19 if non_optional_args:20 return _get_base_type(non_optional_args[0])21 elif (22 getattr(annotation, "__origin__", None) is list23 or getattr(annotation, "__origin__", None) is List24 ):25 assert hasattr(annotation, "__args__") and len(annotation.__args__) >= 1 # type: ignore26 return _get_base_type(annotation.__args__[0]) # type: ignore27 return annotation28 29 30def _contains_list_type(annotation: Type[Any] | None) -> bool:31 origin = getattr(annotation, "__origin__", None)32 33 if origin is list or origin is List:34 return True35 elif origin in (Literal, Union):36 return any(_contains_list_type(arg) for arg in annotation.__args__) # type: ignore37 else:38 return False39 40 41def _parse_bool_arg(arg: str | bytes | bool) -> bool:42 if isinstance(arg, bytes):43 arg = arg.decode("utf-8")44 45 true_values = {"1", "on", "t", "true", "y", "yes"}46 false_values = {"0", "off", "f", "false", "n", "no"}47 48 arg_str = str(arg).lower().strip()49 50 if arg_str in true_values:51 return True52 elif arg_str in false_values:53 return False54 else:55 raise ValueError(f"Invalid boolean argument: {arg}")56 57 58def add_args_from_model(parser: argparse.ArgumentParser, model: Type[BaseModel]):59 """Add arguments from a pydantic model to an argparse parser."""60 61 for name, field in model.model_fields.items():62 description = field.description63 if field.default and description and not field.is_required():64 description += f" (default: {field.default})"65 base_type = (66 _get_base_type(field.annotation) if field.annotation is not None else str67 )68 list_type = _contains_list_type(field.annotation)69 if base_type is not bool:70 parser.add_argument(71 f"--{name}",72 dest=name,73 nargs="*" if list_type else None,74 type=base_type,75 help=description,76 )77 if base_type is bool:78 parser.add_argument(79 f"--{name}",80 dest=name,81 type=_parse_bool_arg,82 help=f"{description}",83 )84 85 86T = TypeVar("T", bound=Type[BaseModel])87 88 89def parse_model_from_args(model: T, args: argparse.Namespace) -> T:90 """Parse a pydantic model from an argparse namespace."""91 return model(92 **{93 k: v94 for k, v in vars(args).items()95 if v is not None and k in model.model_fields96 }97 )98 