Coverage for src/lilbee/catalog/families.py: 100%

32 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-09-17 10:02 +0000

1"""Group model picks into display families.""" 

2 

3import re 

4 

5from lilbee.catalog.formatting import clean_display_name, extract_quant 

6from lilbee.catalog.models import CatalogModel, ModelFamily, ModelVariant 

7from lilbee.catalog.picks import picks_for 

8from lilbee.catalog.types import ModelTask 

9 

10_FAMILY_NAME_RE = re.compile(r"^(.+?)\s+\d") 

11 

12 

13def _extract_family_name(model_name: str) -> str: 

14 """Extract the family name by stripping the trailing parameter count. 

15 Applies clean_display_name first to strip -GGUF, -Instruct, etc. 

16 

17 "Qwen3 8B" -> "Qwen3", "Qwen3-Coder 30B A3B" -> "Qwen3-Coder", 

18 "Nomic Embed Text v1.5" -> "Nomic Embed Text v1.5" (no trailing number pattern). 

19 """ 

20 cleaned = clean_display_name(model_name) 

21 m = _FAMILY_NAME_RE.match(cleaned) 

22 return m.group(1) if m else cleaned 

23 

24 

25def _catalog_to_variant(model: CatalogModel) -> ModelVariant: 

26 """Convert a CatalogModel to a ModelVariant.""" 

27 # Local import to avoid pulling formatting helpers into hf_client/featured. 

28 from lilbee.catalog.formatting import derive_param_count 

29 

30 return ModelVariant( 

31 hf_repo=model.hf_repo, 

32 filename=model.gguf_filename, 

33 param_count=derive_param_count(model), 

34 quant=extract_quant(model.gguf_filename), 

35 size_mb=int(model.size_gb * 1024), 

36 compat=model.compat, 

37 safety_stripped=model.safety_stripped, 

38 ) 

39 

40 

41def _family_slug(display_name: str) -> str: 

42 """Stable slug for a family, derived from its display name.""" 

43 return _extract_family_name(display_name).lower().replace(" ", "-") 

44 

45 

46def _build_families(models: tuple[CatalogModel, ...], task: ModelTask) -> list[ModelFamily]: 

47 """Group CatalogModels into families by display-derived family name.""" 

48 groups: dict[str, list[CatalogModel]] = {} 

49 order: list[str] = [] 

50 for m in models: 

51 family = _extract_family_name(m.display_name) 

52 if family not in groups: 

53 order.append(family) 

54 groups.setdefault(family, []).append(m) 

55 

56 families: list[ModelFamily] = [] 

57 for family_name in order: 

58 members = groups[family_name] 

59 representative = members[0] 

60 variants = [_catalog_to_variant(m) for m in members] 

61 families.append( 

62 ModelFamily( 

63 slug=_family_slug(representative.display_name), 

64 name=family_name, 

65 task=task, 

66 description=representative.description, 

67 variants=tuple(variants), 

68 ) 

69 ) 

70 return families 

71 

72 

73def get_families() -> list[ModelFamily]: 

74 """Get the current picks grouped into families. 

75 Returns families ordered: chat, then embedding, then vision, then reranker. 

76 Within each family, variants preserve the order the picks arrive in, which 

77 is most-popular-first within each parameter tier. 

78 """ 

79 return ( 

80 _build_families(picks_for(ModelTask.CHAT), ModelTask.CHAT) 

81 + _build_families(picks_for(ModelTask.EMBEDDING), ModelTask.EMBEDDING) 

82 + _build_families(picks_for(ModelTask.VISION), ModelTask.VISION) 

83 + _build_families(picks_for(ModelTask.RERANK), ModelTask.RERANK) 

84 )