Team Ai
Apppublic

huggingface/Model_Cards_Writing_Tool

sourceHugging Facemitupdated 2y agoView on Hugging Face
125likes
extract_code.py532 linesDownload Raw Back to root
1#!/usr/bin/env python32 3import re4 5"""6Extracts code from the file "./Libraries.ts".7(Note that "Libraries.ts", must be in the same directory as 8this script).9"""10 11file = None12 13def read_file(library: str, model_name: str) -> str:14    text = file15 16    match = re.search('const ' + library + '.*', text, re.DOTALL).group()17    if match:18        text = match[match.index('`') + 1:match.index('`;')].replace('${model.id}', model_name)19 20    return text21 22file = """23import type { ModelData } from "./Types";24/**25 * Add your new library here.26 */27export enum ModelLibrary {28	"adapter-transformers"   = "Adapter Transformers",29	"allennlp"               = "allenNLP",30	"asteroid"               = "Asteroid",31	"diffusers"              = "Diffusers",32	"espnet"                 = "ESPnet",33	"fairseq"                = "Fairseq",34	"flair"                  = "Flair",35	"keras"                  = "Keras",36	"nemo"                   = "NeMo",37	"pyannote-audio"         = "pyannote.audio",38	"sentence-transformers"  = "Sentence Transformers",39	"sklearn"                = "Scikit-learn",40	"spacy"                  = "spaCy",41	"speechbrain"            = "speechbrain",42	"tensorflowtts"          = "TensorFlowTTS",43	"timm"                   = "Timm",44	"fastai"                 = "fastai",45	"transformers"           = "Transformers",46	"stanza"                 = "Stanza",47	"fasttext"               = "fastText",48	"stable-baselines3"      = "Stable-Baselines3",49	"ml-agents"              = "ML-Agents",50}51 52export const ALL_MODEL_LIBRARY_KEYS = Object.keys(ModelLibrary) as (keyof typeof ModelLibrary)[];53 54 55/**56 * Elements configurable by a model library.57 */58export interface LibraryUiElement {59	/**60	 * Name displayed on the main61	 * call-to-action button on the model page.62	 */63	btnLabel:  string;64	/**65	 * Repo name66	 */67	repoName: string;68	/**69	 * URL to library's repo70	 */71	repoUrl:   string;72	/**73	 * Code snippet displayed on model page74	 */75	snippet:   (model: ModelData) => string;76}77 78function nameWithoutNamespace(modelId: string): string {79	const splitted = modelId.split("/");80	return splitted.length === 1 ? splitted[0] : splitted[1];81}82 83//#region snippets84 85const adapter_transformers = (model: ModelData) =>86	`from transformers import ${model.config?.adapter_transformers?.model_class}87 88model = ${model.config?.adapter_transformers?.model_class}.from_pretrained("${model.config?.adapter_transformers?.{model.id}}")89model.load_adapter("${model.id}", source="hf")`;90 91const allennlpUnknown = (model: ModelData) =>92	`import allennlp_models93from allennlp.predictors.predictor import Predictor94 95predictor = Predictor.from_path("hf://${model.id}")`;96 97const allennlpQuestionAnswering = (model: ModelData) =>98	`import allennlp_models99from allennlp.predictors.predictor import Predictor100 101predictor = Predictor.from_path("hf://${model.id}")102predictor_input = {"passage": "My name is Wolfgang and I live in Berlin", "question": "Where do I live?"}103predictions = predictor.predict_json(predictor_input)`;104 105const allennlp = (model: ModelData) => {106	if (model.tags?.includes("question-answering")) {107		return allennlpQuestionAnswering(model);108	}109	return allennlpUnknown(model);110};111 112const asteroid = (model: ModelData) =>113	`from asteroid.models import BaseModel114 115model = BaseModel.from_pretrained("${model.id}")`;116 117const diffusers = (model: ModelData) =>118	`from diffusers import DiffusionPipeline119 120pipeline = DiffusionPipeline.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`;121 122const espnetTTS = (model: ModelData) =>123	`from espnet2.bin.tts_inference import Text2Speech124 125model = Text2Speech.from_pretrained("${model.id}")126 127speech, *_ = model("text to generate speech from")`;128 129const espnetASR = (model: ModelData) =>130	`from espnet2.bin.asr_inference import Speech2Text131 132model = Speech2Text.from_pretrained(133  "${model.id}"134)135 136speech, rate = soundfile.read("speech.wav")137text, *_ = model(speech)`;138 139const espnetUnknown = () =>140	`unknown model type (must be text-to-speech or automatic-speech-recognition)`;141 142const espnet = (model: ModelData) => {143	if (model.tags?.includes("text-to-speech")) {144		return espnetTTS(model);145	} else if (model.tags?.includes("automatic-speech-recognition")) {146		return espnetASR(model);147	}148	return espnetUnknown();149};150 151const fairseq = (model: ModelData) =>152	`from fairseq.checkpoint_utils import load_model_ensemble_and_task_from_hf_hub153 154models, cfg, task = load_model_ensemble_and_task_from_hf_hub(155    "${model.id}"156)`;157 158 159const flair = (model: ModelData) =>160	`from flair.models import SequenceTagger161 162tagger = SequenceTagger.load("${model.id}")`;163 164const keras = (model: ModelData) =>165	`from huggingface_hub import from_pretrained_keras166 167model = from_pretrained_keras("${model.id}")168`;169 170const pyannote_audio_pipeline = (model: ModelData) =>171	`from pyannote.audio import Pipeline172  173pipeline = Pipeline.from_pretrained("${model.id}")174 175# inference on the whole file176pipeline("file.wav")177 178# inference on an excerpt179from pyannote.core import Segment180excerpt = Segment(start=2.0, end=5.0)181 182from pyannote.audio import Audio183waveform, sample_rate = Audio().crop("file.wav", excerpt)184pipeline({"waveform": waveform, "sample_rate": sample_rate})`;185 186const pyannote_audio_model = (model: ModelData) =>187	`from pyannote.audio import Model, Inference188 189model = Model.from_pretrained("${model.id}")190inference = Inference(model)191 192# inference on the whole file193inference("file.wav")194 195# inference on an excerpt196from pyannote.core import Segment197excerpt = Segment(start=2.0, end=5.0)198inference.crop("file.wav", excerpt)`;199 200const pyannote_audio = (model: ModelData) => {201	if (model.tags?.includes("pyannote-audio-pipeline")) {202		return pyannote_audio_pipeline(model);203	}204	return pyannote_audio_model(model);205};206 207const tensorflowttsTextToMel = (model: ModelData) =>208	`from tensorflow_tts.inference import AutoProcessor, TFAutoModel209 210processor = AutoProcessor.from_pretrained("${model.id}")211model = TFAutoModel.from_pretrained("${model.id}")212`;213 214const tensorflowttsMelToWav = (model: ModelData) =>215	`from tensorflow_tts.inference import TFAutoModel216 217model = TFAutoModel.from_pretrained("${model.id}")218audios = model.inference(mels)219`;220 221const tensorflowttsUnknown = (model: ModelData) =>222	`from tensorflow_tts.inference import TFAutoModel223 224model = TFAutoModel.from_pretrained("${model.id}")225`;226 227const tensorflowtts = (model: ModelData) => {228	if (model.tags?.includes("text-to-mel")) {229		return tensorflowttsTextToMel(model);230	} else if (model.tags?.includes("mel-to-wav")) {231		return tensorflowttsMelToWav(model);232	}233	return tensorflowttsUnknown(model);234};235 236const timm = (model: ModelData) =>237	`import timm238 239model = timm.create_model("hf_hub:${model.id}", pretrained=True)`;240 241const sklearn = (model: ModelData) =>242	`from huggingface_hub import hf_hub_download243import joblib244 245model = joblib.load(246	hf_hub_download("${model.id}", "sklearn_model.joblib")247)`;248 249const fastai = (model: ModelData) =>250	`from huggingface_hub import from_pretrained_fastai251 252learn = from_pretrained_fastai("${model.id}")`;253 254const sentenceTransformers = (model: ModelData) =>255	`from sentence_transformers import SentenceTransformer256 257model = SentenceTransformer("${model.id}")`;258 259const spacy = (model: ModelData) =>260	`!pip install https://huggingface.co/${model.id}/resolve/main/${nameWithoutNamespace(model.id)}-any-py3-none-any.whl261 262# Using spacy.load().263import spacy264nlp = spacy.load("${nameWithoutNamespace(model.id)}")265 266# Importing as module.267import ${nameWithoutNamespace(model.id)}268nlp = ${nameWithoutNamespace(model.id)}.load()`;269 270const stanza = (model: ModelData) =>271	`import stanza272 273stanza.download("${nameWithoutNamespace(model.id).replace("stanza-", "")}")274nlp = stanza.Pipeline("${nameWithoutNamespace(model.id).replace("stanza-", "")}")`;275 276 277const speechBrainMethod = (speechbrainInterface: string) => {278	switch (speechbrainInterface) {279		case "EncoderClassifier":280		   return "classify_file";281		case "EncoderDecoderASR":282		case "EncoderASR":283			return "transcribe_file";284		case "SpectralMaskEnhancement":285			return "enhance_file";286		case "SepformerSeparation":287			return "separate_file";288		default:289			return undefined;290	}291};292 293const speechbrain = (model: ModelData) => {294	const speechbrainInterface = model.config?.speechbrain?.interface;295	if (speechbrainInterface === undefined) {296		return `# interface not specified in config.json`;297	}298 299	const speechbrainMethod = speechBrainMethod(speechbrainInterface);300	if (speechbrainMethod === undefined) {301		return `# interface in config.json invalid`;302	}303 304	return `from speechbrain.pretrained import ${speechbrainInterface}305model = ${speechbrainInterface}.from_hparams(306  "${model.id}"307)308model.${speechbrainMethod}("file.wav")`;309};310 311const transformers = (model: ModelData) => {312	const info = model.transformersInfo;313	if (!info) {314		return `# ⚠️ Type of model unknown`;315	}316	if (info.processor) {317		const varName = info.processor === "AutoTokenizer" ? "tokenizer"318			: info.processor === "AutoFeatureExtractor" ? "extractor"319				: "processor"320		;321		return [322			`from transformers import ${info.processor}, ${info.auto_model}`,323			"",324			`${varName} = ${info.processor}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,325			"",326			`model = ${info.auto_model}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,327		].join("\n");328	} else {329		return [330			`from transformers import ${info.auto_model}`,331			"",332			`model = ${info.auto_model}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,333		].join("\n");334	}335};336 337const fasttext = (model: ModelData) =>338	`from huggingface_hub import hf_hub_download339import fasttext340 341model = fasttext.load_model(hf_hub_download("${model.id}", "model.bin"))`;342 343const stableBaselines3 = (model: ModelData) =>344	`from huggingface_sb3 import load_from_hub345checkpoint = load_from_hub(346	repo_id="${model.id}",347	filename="{MODEL FILENAME}.zip",348)`;349 350const nemoDomainResolver = (domain: string, model: ModelData): string | undefined => {351	const modelName = `${nameWithoutNamespace(model.id)}.nemo`;352 353	switch (domain) {354		case "ASR":355			return `import nemo.collections.asr as nemo_asr356asr_model = nemo_asr.models.ASRModel.from_pretrained("${model.id}")357 358transcriptions = asr_model.transcribe(["file.wav"])`;359		default:360			return undefined;361	}362};363 364const mlAgents = (model: ModelData) =>365	`mlagents-load-from-hf --repo-id="${model.id}" --local-dir="./downloads"`;366	367const nemo = (model: ModelData) => {368	let command: string | undefined = undefined;369	// Resolve the tag to a nemo domain/sub-domain 370	if (model.tags?.includes("automatic-speech-recognition")) {371		command = nemoDomainResolver("ASR", model);372	}373	374	return command ?? `# tag did not correspond to a valid NeMo domain.`;375};376 377//#endregion378 379 380 381export const MODEL_LIBRARIES_UI_ELEMENTS: { [key in keyof typeof ModelLibrary]?: LibraryUiElement } = {382	// ^^ TODO(remove the optional ? marker when Stanza snippet is available)383	"adapter-transformers": {384		btnLabel: "Adapter Transformers",385		repoName: "adapter-transformers",386		repoUrl:  "https://github.com/Adapter-Hub/adapter-transformers",387		snippet:  adapter_transformers,388	},389	"allennlp": {390		btnLabel: "AllenNLP",391		repoName: "AllenNLP",392		repoUrl:  "https://github.com/allenai/allennlp",393		snippet:  allennlp,394	},395	"asteroid": {396		btnLabel: "Asteroid",397		repoName: "Asteroid",398		repoUrl:  "https://github.com/asteroid-team/asteroid",399		snippet:  asteroid,400	},401	"diffusers": {402		btnLabel: "Diffusers",403		repoName: "🤗/diffusers",404		repoUrl:  "https://github.com/huggingface/diffusers",405		snippet:  diffusers,406	},407	"espnet": {408		btnLabel: "ESPnet",409		repoName: "ESPnet",410		repoUrl:  "https://github.com/espnet/espnet",411		snippet:  espnet,412	},413	"fairseq": {414		btnLabel: "Fairseq",415		repoName: "fairseq",416		repoUrl:  "https://github.com/pytorch/fairseq",417		snippet:  fairseq,418	},419	"flair": {420		btnLabel: "Flair",421		repoName: "Flair",422		repoUrl:  "https://github.com/flairNLP/flair",423		snippet:  flair,424	},425	"keras": {426		btnLabel: "Keras",427		repoName: "Keras",428		repoUrl:  "https://github.com/keras-team/keras",429		snippet:  keras,430	},431	"nemo": {432		btnLabel: "NeMo",433		repoName: "NeMo",434		repoUrl:  "https://github.com/NVIDIA/NeMo",435		snippet:  nemo,436	},437	"pyannote-audio": {438		btnLabel: "pyannote.audio",439		repoName: "pyannote-audio",440		repoUrl:  "https://github.com/pyannote/pyannote-audio",441		snippet:  pyannote_audio,442	},443	"sentence-transformers": {444		btnLabel: "sentence-transformers",445		repoName: "sentence-transformers",446		repoUrl:  "https://github.com/UKPLab/sentence-transformers",447		snippet:  sentenceTransformers,448	},449	"sklearn": {450		btnLabel: "Scikit-learn",451		repoName: "Scikit-learn",452		repoUrl:  "https://github.com/scikit-learn/scikit-learn",453		snippet:  sklearn,454	},455	"fastai": {456		btnLabel: "fastai",457		repoName: "fastai",458		repoUrl:  "https://github.com/fastai/fastai",459		snippet:  fastai,460	},461	"spacy": {462		btnLabel: "spaCy",463		repoName: "spaCy",464		repoUrl:  "https://github.com/explosion/spaCy",465		snippet:  spacy,466	},467	"speechbrain": {468		btnLabel: "speechbrain",469		repoName: "speechbrain",470		repoUrl:  "https://github.com/speechbrain/speechbrain",471		snippet:  speechbrain,472	},473	"stanza": {474		btnLabel: "Stanza",475		repoName: "stanza",476		repoUrl: "https://github.com/stanfordnlp/stanza",477		snippet: stanza,478	},479	"tensorflowtts": {480		btnLabel: "TensorFlowTTS",481		repoName: "TensorFlowTTS",482		repoUrl:  "https://github.com/TensorSpeech/TensorFlowTTS",483		snippet:  tensorflowtts,484	},485	"timm": {486		btnLabel: "timm",487		repoName: "pytorch-image-models",488		repoUrl:  "https://github.com/rwightman/pytorch-image-models",489		snippet:  timm,490	},491	"transformers": {492		btnLabel: "Transformers",493		repoName: "🤗/transformers",494		repoUrl:  "https://github.com/huggingface/transformers",495		snippet:  transformers,496	},497	"fasttext": {498		btnLabel: "fastText",499		repoName: "fastText",500		repoUrl:  "https://fasttext.cc/",501		snippet:  fasttext,502	},503	"stable-baselines3": {504		btnLabel: "stable-baselines3",505		repoName: "stable-baselines3",506		repoUrl:  "https://github.com/huggingface/huggingface_sb3",507		snippet:  stableBaselines3,508	},509	"ml-agents": {510		btnLabel: "ml-agents",511		repoName: "ml-agents",512		repoUrl:  "https://github.com/huggingface/ml-agents",513		snippet:  mlAgents,514	},515} as const;516"""517 518 519if __name__ == '__main__':520    import sys521    library_name = "keras"522    model_name = "Distillgpt2"523    print(read_file(library_name, model_name))524    525    """"526    try:527        args = sys.argv[1:]528        if args:529            print(read_file(args[0], args[1]))530    except IndexError:531        pass532    """