LLM-XRay / src /display /utils.py
jmullings
Base application
7774431
Raw
History Blame Contribute Delete
3.3 kB
from dataclasses import dataclass, make_dataclass
from enum import Enum
try:
from src.about import Tasks
except ImportError:
from about import Tasks
@dataclass(frozen=True)
class ColumnContent:
name: str
type: str
displayed_by_default: bool
hidden: bool = False
never_hidden: bool = False
def make_auto_eval_column_dict():
cols = []
# "T" column set to "markdown" so it can be a clickable link
cols.append(["model_type_symbol", ColumnContent, ColumnContent("T", "markdown", True, never_hidden=True)])
cols.append(["model", ColumnContent, ColumnContent("Model", "markdown", True, never_hidden=True)])
for task in Tasks:
cols.append([task.name, ColumnContent, ColumnContent(task.value.col_name, "number", True)])
cols.append(["sample_adequate", ColumnContent, ColumnContent("Calib. Sample OK", "str", True)])
cols.append(["architecture", ColumnContent, ColumnContent("Architecture", "str", False)])
cols.append(["precision", ColumnContent, ColumnContent("Precision", "str", False)])
cols.append(["params", ColumnContent, ColumnContent("#Params (B)", "number", True)])
cols.append(["license", ColumnContent, ColumnContent("Hub License", "str", False)])
cols.append(["revision", ColumnContent, ColumnContent("Model sha", "str", False)])
return cols
auto_eval_column_dict = make_auto_eval_column_dict()
AutoEvalColumn = make_dataclass("AutoEvalColumn", auto_eval_column_dict, frozen=True)
def fields(raw_class=None):
return [col[2] for col in auto_eval_column_dict]
@dataclass
class ModelDetails:
name: str
display_name: str = ""
symbol: str = ""
class ModelType(Enum):
PT = ModelDetails(name="pretrained", symbol="🟒")
FT = ModelDetails(name="fine-tuned", symbol="πŸ”Ά")
IFT = ModelDetails(name="instruction-tuned", symbol="β­•")
RL = ModelDetails(name="RL-tuned", symbol="🟦")
Unknown = ModelDetails(name="", symbol="?")
def to_str(self, separator=" "):
return f"{self.value.symbol}{separator}{self.value.name}"
@staticmethod
def from_str(type_str):
if not type_str:
return ModelType.Unknown
if "fine-tuned" in type_str or "πŸ”Ά" in type_str:
return ModelType.FT
if "pretrained" in type_str or "🟒" in type_str:
return ModelType.PT
if "RL-tuned" in type_str or "🟦" in type_str:
return ModelType.RL
if "instruction-tuned" in type_str or "β­•" in type_str:
return ModelType.IFT
return ModelType.Unknown
class Precision(Enum):
bfloat16 = ModelDetails("bfloat16")
float16 = ModelDetails("float16")
float32 = ModelDetails("float32")
Unknown = ModelDetails("?")
@staticmethod
def from_str(prec):
if not prec:
return Precision.Unknown
prec_str = str(prec).lower()
if "bfloat16" in prec_str:
return Precision.bfloat16
if "float16" in prec_str:
return Precision.float16
if "float32" in prec_str:
return Precision.float32
return Precision.Unknown
COLS = [c.name for c in fields()]
EVAL_COLS = [c.name for c in fields()]
EVAL_TYPES = [c.type for c in fields()]
BENCHMARK_COLS = [t.value.col_name for t in Tasks]