Team Ai
Apppublic

huggingface/Model_Cards_Writing_Tool

sourceHugging Facemitupdated 2y agoView on Hugging Face
125likes
extract_code.cpython-39.pyc500 linesDownload Raw Back to __pycache__
1a

2W�*c�8�@sLddlZdZeeed�dd�ZdZedkrHddlZdZdZeeee��dS)	�N)�library�3model_name�returncCsLt}t�d|d|tj���}|rH||�d�d|�d���d|�}|S)Nzconst z.*�`�z`;z${model.id})�file�re�search�DOTALL�group�index�replace)rr�text�match�r�M/Users/ezi/Desktop/HF/Prompt_writing/ModelCard_Prompt_Writing/extract_code.py�	read_file
s4$ru�55import type { ModelData } from "./Types";6/**7 * Add your new library here.8 */9export enum ModelLibrary {10	"adapter-transformers"   = "Adapter Transformers",11	"allennlp"               = "allenNLP",12	"asteroid"               = "Asteroid",13	"diffusers"              = "Diffusers",14	"espnet"                 = "ESPnet",15	"fairseq"                = "Fairseq",16	"flair"                  = "Flair",17	"keras"                  = "Keras",18	"nemo"                   = "NeMo",19	"pyannote-audio"         = "pyannote.audio",20	"sentence-transformers"  = "Sentence Transformers",21	"sklearn"                = "Scikit-learn",22	"spacy"                  = "spaCy",23	"speechbrain"            = "speechbrain",24	"tensorflowtts"          = "TensorFlowTTS",25	"timm"                   = "Timm",26	"fastai"                 = "fastai",27	"transformers"           = "Transformers",28	"stanza"                 = "Stanza",29	"fasttext"               = "fastText",30	"stable-baselines3"      = "Stable-Baselines3",31	"ml-agents"              = "ML-Agents",32}33 34export const ALL_MODEL_LIBRARY_KEYS = Object.keys(ModelLibrary) as (keyof typeof ModelLibrary)[];35 36 37/**38 * Elements configurable by a model library.39 */40export interface LibraryUiElement {41	/**42	 * Name displayed on the main43	 * call-to-action button on the model page.44	 */45	btnLabel:  string;46	/**47	 * Repo name48	 */49	repoName: string;50	/**51	 * URL to library's repo52	 */53	repoUrl:   string;54	/**55	 * Code snippet displayed on model page56	 */57	snippet:   (model: ModelData) => string;58}59 60function nameWithoutNamespace(modelId: string): string {61	const splitted = modelId.split("/");62	return splitted.length === 1 ? splitted[0] : splitted[1];63}64 65//#region snippets66 67const adapter_transformers = (model: ModelData) =>68	`from transformers import ${model.config?.adapter_transformers?.model_class}69 70model = ${model.config?.adapter_transformers?.model_class}.from_pretrained("${model.config?.adapter_transformers?.{model.id}}")71model.load_adapter("${model.id}", source="hf")`;72 73const allennlpUnknown = (model: ModelData) =>74	`import allennlp_models75from allennlp.predictors.predictor import Predictor76 77predictor = Predictor.from_path("hf://${model.id}")`;78 79const allennlpQuestionAnswering = (model: ModelData) =>80	`import allennlp_models81from allennlp.predictors.predictor import Predictor82 83predictor = Predictor.from_path("hf://${model.id}")84predictor_input = {"passage": "My name is Wolfgang and I live in Berlin", "question": "Where do I live?"}85predictions = predictor.predict_json(predictor_input)`;86 87const allennlp = (model: ModelData) => {88	if (model.tags?.includes("question-answering")) {89		return allennlpQuestionAnswering(model);90	}91	return allennlpUnknown(model);92};93 94const asteroid = (model: ModelData) =>95	`from asteroid.models import BaseModel96 97model = BaseModel.from_pretrained("${model.id}")`;98 99const diffusers = (model: ModelData) =>100	`from diffusers import DiffusionPipeline101 102pipeline = DiffusionPipeline.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`;103 104const espnetTTS = (model: ModelData) =>105	`from espnet2.bin.tts_inference import Text2Speech106 107model = Text2Speech.from_pretrained("${model.id}")108 109speech, *_ = model("text to generate speech from")`;110 111const espnetASR = (model: ModelData) =>112	`from espnet2.bin.asr_inference import Speech2Text113 114model = Speech2Text.from_pretrained(115  "${model.id}"116)117 118speech, rate = soundfile.read("speech.wav")119text, *_ = model(speech)`;120 121const espnetUnknown = () =>122	`unknown model type (must be text-to-speech or automatic-speech-recognition)`;123 124const espnet = (model: ModelData) => {125	if (model.tags?.includes("text-to-speech")) {126		return espnetTTS(model);127	} else if (model.tags?.includes("automatic-speech-recognition")) {128		return espnetASR(model);129	}130	return espnetUnknown();131};132 133const fairseq = (model: ModelData) =>134	`from fairseq.checkpoint_utils import load_model_ensemble_and_task_from_hf_hub135 136models, cfg, task = load_model_ensemble_and_task_from_hf_hub(137    "${model.id}"138)`;139 140 141const flair = (model: ModelData) =>142	`from flair.models import SequenceTagger143 144tagger = SequenceTagger.load("${model.id}")`;145 146const keras = (model: ModelData) =>147	`from huggingface_hub import from_pretrained_keras148 149model = from_pretrained_keras("${model.id}")150`;151 152const pyannote_audio_pipeline = (model: ModelData) =>153	`from pyannote.audio import Pipeline154  155pipeline = Pipeline.from_pretrained("${model.id}")156 157# inference on the whole file158pipeline("file.wav")159 160# inference on an excerpt161from pyannote.core import Segment162excerpt = Segment(start=2.0, end=5.0)163 164from pyannote.audio import Audio165waveform, sample_rate = Audio().crop("file.wav", excerpt)166pipeline({"waveform": waveform, "sample_rate": sample_rate})`;167 168const pyannote_audio_model = (model: ModelData) =>169	`from pyannote.audio import Model, Inference170 171model = Model.from_pretrained("${model.id}")172inference = Inference(model)173 174# inference on the whole file175inference("file.wav")176 177# inference on an excerpt178from pyannote.core import Segment179excerpt = Segment(start=2.0, end=5.0)180inference.crop("file.wav", excerpt)`;181 182const pyannote_audio = (model: ModelData) => {183	if (model.tags?.includes("pyannote-audio-pipeline")) {184		return pyannote_audio_pipeline(model);185	}186	return pyannote_audio_model(model);187};188 189const tensorflowttsTextToMel = (model: ModelData) =>190	`from tensorflow_tts.inference import AutoProcessor, TFAutoModel191 192processor = AutoProcessor.from_pretrained("${model.id}")193model = TFAutoModel.from_pretrained("${model.id}")194`;195 196const tensorflowttsMelToWav = (model: ModelData) =>197	`from tensorflow_tts.inference import TFAutoModel198 199model = TFAutoModel.from_pretrained("${model.id}")200audios = model.inference(mels)201`;202 203const tensorflowttsUnknown = (model: ModelData) =>204	`from tensorflow_tts.inference import TFAutoModel205 206model = TFAutoModel.from_pretrained("${model.id}")207`;208 209const tensorflowtts = (model: ModelData) => {210	if (model.tags?.includes("text-to-mel")) {211		return tensorflowttsTextToMel(model);212	} else if (model.tags?.includes("mel-to-wav")) {213		return tensorflowttsMelToWav(model);214	}215	return tensorflowttsUnknown(model);216};217 218const timm = (model: ModelData) =>219	`import timm220 221model = timm.create_model("hf_hub:${model.id}", pretrained=True)`;222 223const sklearn = (model: ModelData) =>224	`from huggingface_hub import hf_hub_download225import joblib226 227model = joblib.load(228	hf_hub_download("${model.id}", "sklearn_model.joblib")229)`;230 231const fastai = (model: ModelData) =>232	`from huggingface_hub import from_pretrained_fastai233 234learn = from_pretrained_fastai("${model.id}")`;235 236const sentenceTransformers = (model: ModelData) =>237	`from sentence_transformers import SentenceTransformer238 239model = SentenceTransformer("${model.id}")`;240 241const spacy = (model: ModelData) =>242	`!pip install https://huggingface.co/${model.id}/resolve/main/${nameWithoutNamespace(model.id)}-any-py3-none-any.whl243 244# Using spacy.load().245import spacy246nlp = spacy.load("${nameWithoutNamespace(model.id)}")247 248# Importing as module.249import ${nameWithoutNamespace(model.id)}250nlp = ${nameWithoutNamespace(model.id)}.load()`;251 252const stanza = (model: ModelData) =>253	`import stanza254 255stanza.download("${nameWithoutNamespace(model.id).replace("stanza-", "")}")256nlp = stanza.Pipeline("${nameWithoutNamespace(model.id).replace("stanza-", "")}")`;257 258 259const speechBrainMethod = (speechbrainInterface: string) => {260	switch (speechbrainInterface) {261		case "EncoderClassifier":262		   return "classify_file";263		case "EncoderDecoderASR":264		case "EncoderASR":265			return "transcribe_file";266		case "SpectralMaskEnhancement":267			return "enhance_file";268		case "SepformerSeparation":269			return "separate_file";270		default:271			return undefined;272	}273};274 275const speechbrain = (model: ModelData) => {276	const speechbrainInterface = model.config?.speechbrain?.interface;277	if (speechbrainInterface === undefined) {278		return `# interface not specified in config.json`;279	}280 281	const speechbrainMethod = speechBrainMethod(speechbrainInterface);282	if (speechbrainMethod === undefined) {283		return `# interface in config.json invalid`;284	}285 286	return `from speechbrain.pretrained import ${speechbrainInterface}287model = ${speechbrainInterface}.from_hparams(288  "${model.id}"289)290model.${speechbrainMethod}("file.wav")`;291};292 293const transformers = (model: ModelData) => {294	const info = model.transformersInfo;295	if (!info) {296		return `# ⚠️ Type of model unknown`;297	}298	if (info.processor) {299		const varName = info.processor === "AutoTokenizer" ? "tokenizer"300			: info.processor === "AutoFeatureExtractor" ? "extractor"301				: "processor"302		;303		return [304			`from transformers import ${info.processor}, ${info.auto_model}`,305			"",306			`${varName} = ${info.processor}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,307			"",308			`model = ${info.auto_model}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,309		].join("310");311	} else {312		return [313			`from transformers import ${info.auto_model}`,314			"",315			`model = ${info.auto_model}.from_pretrained("${model.id}"${model.private ? ", use_auth_token=True" : ""})`,316		].join("317");318	}319};320 321const fasttext = (model: ModelData) =>322	`from huggingface_hub import hf_hub_download323import fasttext324 325model = fasttext.load_model(hf_hub_download("${model.id}", "model.bin"))`;326 327const stableBaselines3 = (model: ModelData) =>328	`from huggingface_sb3 import load_from_hub329checkpoint = load_from_hub(330	repo_id="${model.id}",331	filename="{MODEL FILENAME}.zip",332)`;333 334const nemoDomainResolver = (domain: string, model: ModelData): string | undefined => {335	const modelName = `${nameWithoutNamespace(model.id)}.nemo`;336 337	switch (domain) {338		case "ASR":339			return `import nemo.collections.asr as nemo_asr340asr_model = nemo_asr.models.ASRModel.from_pretrained("${model.id}")341 342transcriptions = asr_model.transcribe(["file.wav"])`;343		default:344			return undefined;345	}346};347 348const mlAgents = (model: ModelData) =>349	`mlagents-load-from-hf --repo-id="${model.id}" --local-dir="./downloads"`;350	351const nemo = (model: ModelData) => {352	let command: string | undefined = undefined;353	// Resolve the tag to a nemo domain/sub-domain 354	if (model.tags?.includes("automatic-speech-recognition")) {355		command = nemoDomainResolver("ASR", model);356	}357	358	return command ?? `# tag did not correspond to a valid NeMo domain.`;359};360 361//#endregion362 363 364 365export const MODEL_LIBRARIES_UI_ELEMENTS: { [key in keyof typeof ModelLibrary]?: LibraryUiElement } = {366	// ^^ TODO(remove the optional ? marker when Stanza snippet is available)367	"adapter-transformers": {368		btnLabel: "Adapter Transformers",369		repoName: "adapter-transformers",370		repoUrl:  "https://github.com/Adapter-Hub/adapter-transformers",371		snippet:  adapter_transformers,372	},373	"allennlp": {374		btnLabel: "AllenNLP",375		repoName: "AllenNLP",376		repoUrl:  "https://github.com/allenai/allennlp",377		snippet:  allennlp,378	},379	"asteroid": {380		btnLabel: "Asteroid",381		repoName: "Asteroid",382		repoUrl:  "https://github.com/asteroid-team/asteroid",383		snippet:  asteroid,384	},385	"diffusers": {386		btnLabel: "Diffusers",387		repoName: "🤗/diffusers",388		repoUrl:  "https://github.com/huggingface/diffusers",389		snippet:  diffusers,390	},391	"espnet": {392		btnLabel: "ESPnet",393		repoName: "ESPnet",394		repoUrl:  "https://github.com/espnet/espnet",395		snippet:  espnet,396	},397	"fairseq": {398		btnLabel: "Fairseq",399		repoName: "fairseq",400		repoUrl:  "https://github.com/pytorch/fairseq",401		snippet:  fairseq,402	},403	"flair": {404		btnLabel: "Flair",405		repoName: "Flair",406		repoUrl:  "https://github.com/flairNLP/flair",407		snippet:  flair,408	},409	"keras": {410		btnLabel: "Keras",411		repoName: "Keras",412		repoUrl:  "https://github.com/keras-team/keras",413		snippet:  keras,414	},415	"nemo": {416		btnLabel: "NeMo",417		repoName: "NeMo",418		repoUrl:  "https://github.com/NVIDIA/NeMo",419		snippet:  nemo,420	},421	"pyannote-audio": {422		btnLabel: "pyannote.audio",423		repoName: "pyannote-audio",424		repoUrl:  "https://github.com/pyannote/pyannote-audio",425		snippet:  pyannote_audio,426	},427	"sentence-transformers": {428		btnLabel: "sentence-transformers",429		repoName: "sentence-transformers",430		repoUrl:  "https://github.com/UKPLab/sentence-transformers",431		snippet:  sentenceTransformers,432	},433	"sklearn": {434		btnLabel: "Scikit-learn",435		repoName: "Scikit-learn",436		repoUrl:  "https://github.com/scikit-learn/scikit-learn",437		snippet:  sklearn,438	},439	"fastai": {440		btnLabel: "fastai",441		repoName: "fastai",442		repoUrl:  "https://github.com/fastai/fastai",443		snippet:  fastai,444	},445	"spacy": {446		btnLabel: "spaCy",447		repoName: "spaCy",448		repoUrl:  "https://github.com/explosion/spaCy",449		snippet:  spacy,450	},451	"speechbrain": {452		btnLabel: "speechbrain",453		repoName: "speechbrain",454		repoUrl:  "https://github.com/speechbrain/speechbrain",455		snippet:  speechbrain,456	},457	"stanza": {458		btnLabel: "Stanza",459		repoName: "stanza",460		repoUrl: "https://github.com/stanfordnlp/stanza",461		snippet: stanza,462	},463	"tensorflowtts": {464		btnLabel: "TensorFlowTTS",465		repoName: "TensorFlowTTS",466		repoUrl:  "https://github.com/TensorSpeech/TensorFlowTTS",467		snippet:  tensorflowtts,468	},469	"timm": {470		btnLabel: "timm",471		repoName: "pytorch-image-models",472		repoUrl:  "https://github.com/rwightman/pytorch-image-models",473		snippet:  timm,474	},475	"transformers": {476		btnLabel: "Transformers",477		repoName: "🤗/transformers",478		repoUrl:  "https://github.com/huggingface/transformers",479		snippet:  transformers,480	},481	"fasttext": {482		btnLabel: "fastText",483		repoName: "fastText",484		repoUrl:  "https://fasttext.cc/",485		snippet:  fasttext,486	},487	"stable-baselines3": {488		btnLabel: "stable-baselines3",489		repoName: "stable-baselines3",490		repoUrl:  "https://github.com/huggingface/huggingface_sb3",491		snippet:  stableBaselines3,492	},493	"ml-agents": {494		btnLabel: "ml-agents",495		repoName: "ml-agents",496		repoUrl:  "https://github.com/huggingface/ml-agents",497		snippet:  mlAgents,498	},499} as const;500�__main__�kerasZDistillgpt2)	rr�strr�__name__�sys�library_namer�printrrrr�<module>s	t