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

32 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-09-04 17:08 +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 ) 

38 

39 

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

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

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

43 

44 

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

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

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

48 order: list[str] = [] 

49 for m in models: 

50 family = _extract_family_name(m.display_name) 

51 if family not in groups: 

52 order.append(family) 

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

54 

55 families: list[ModelFamily] = [] 

56 for family_name in order: 

57 members = groups[family_name] 

58 representative = members[0] 

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

60 families.append( 

61 ModelFamily( 

62 slug=_family_slug(representative.display_name), 

63 name=family_name, 

64 task=task, 

65 description=representative.description, 

66 variants=tuple(variants), 

67 ) 

68 ) 

69 return families 

70 

71 

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

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

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

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

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

77 """ 

78 return ( 

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

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

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

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

83 )