m-elio commited on
Commit
9c83fd8
·
verified ·
1 Parent(s): 3ca0380

Create README.md

Browse files
Files changed (1) hide show
  1. README.md +268 -0
README.md ADDED
@@ -0,0 +1,268 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ ---
3
+
4
+ # Model Card for Mimir-1.6B-Instruct
5
+
6
+ ## How to Get Started with the Model
7
+
8
+ Below you can find an example of model usage. To facilitate its usage, we recommend to follow these steps:
9
+
10
+ ```
11
+ huggingface-cli download mimir-lcm/Mimir-1.6B-Instruct --local-dir mimir-lcm/Mimir-1.6B-Instruct
12
+ git clone https://github.com/facebookresearch/large_concept_model.git
13
+ mv large_concept_model/lcm .
14
+ pip install torch==2.5.1 --extra-index-url https://download.pytorch.org/whl/cu121 --upgrade
15
+ pip install fairseq2==v0.3.0rc1 --pre --extra-index-url https://fair.pkg.atmeta.com/fairseq2/whl/rc/pt2.5.1/cu121 --upgrade
16
+ pip install omegaconf==2.3.0
17
+ pip install sonar-space==0.3.2
18
+ pip install wtpsplit==2.1.2
19
+ ```
20
+
21
+ Now you should be able to run the following:
22
+
23
+ ```python
24
+ import lcm
25
+ import torch
26
+ from pathlib import Path
27
+
28
+ from lcm.models.two_tower_diffusion_lcm.builder import (
29
+ create_two_tower_diffusion_lcm_model,
30
+ )
31
+ from lcm.models.two_tower_diffusion_lcm.archs import two_tower_diffusion_lcm_1_6B
32
+ from lcm.inference.two_tower_diffusion_lcm.generator import (
33
+ TwoTowerDiffusionLCMGenerator,
34
+ DiffusionLCMGeneratorOptions,
35
+ )
36
+ from lcm.datasets.batch import EmbeddingsBatch
37
+ from sonar.inference_pipelines.text import TextToEmbeddingModelPipeline, EmbeddingToTextModelPipeline
38
+
39
+ from wtpsplit import SaT
40
+
41
+ lcm.setup_fairseq2()
42
+
43
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
44
+
45
+ from lcm.models.two_tower_diffusion_lcm.builder import TwoTowerDiffusionLCModel
46
+
47
+ _original_sample_fn = TwoTowerDiffusionLCModel.sample_initial_noise_vectors
48
+
49
+ def _patched_sample_fn(self, batch_size: int):
50
+ latents = _original_sample_fn(self, batch_size)
51
+ return latents.to(dtype=self.dtype)
52
+
53
+ TwoTowerDiffusionLCModel.sample_initial_noise_vectors = _patched_sample_fn
54
+
55
+ CHECKPOINT_PATH = "mimir-lcm/Mimir-1.6B-Instruct/model.pt"
56
+ INFERENCE_DTYPE = torch.float16
57
+
58
+ TEXT_DECODER = EmbeddingToTextModelPipeline(decoder="text_sonar_basic_decoder", tokenizer="text_sonar_basic_decoder", device=torch.device(DEVICE))
59
+ TEXT_EMBEDDER = TextToEmbeddingModelPipeline(encoder="text_sonar_basic_encoder", tokenizer="text_sonar_basic_encoder", device=torch.device(DEVICE))
60
+
61
+ def decode_embeddings(embeddings):
62
+
63
+ embeddings = embeddings.to(device=DEVICE, dtype=torch.float32)
64
+
65
+ print("Decoding...")
66
+ results = TEXT_DECODER.predict(
67
+ embeddings,
68
+ target_lang="eng_Latn"
69
+ )
70
+
71
+ return results
72
+
73
+ def get_eos_vector():
74
+ return TEXT_EMBEDDER.predict(["End of text."], source_lang="eng_Latn").squeeze().to(device=DEVICE, dtype=INFERENCE_DTYPE)
75
+
76
+ def load_two_tower_model(checkpoint_path, device="cuda"):
77
+
78
+ config = two_tower_diffusion_lcm_1_6B()
79
+
80
+ print("Building model structure...")
81
+ model = create_two_tower_diffusion_lcm_model(
82
+ config,
83
+ device=torch.device(device),
84
+ dtype=INFERENCE_DTYPE
85
+ )
86
+
87
+ print(f"Loading weights from {checkpoint_path}...")
88
+ state_dict = torch.load(checkpoint_path, map_location=device)
89
+
90
+ if "model" in state_dict:
91
+ state_dict = state_dict["model"]
92
+
93
+ model.load_state_dict(state_dict, strict=True)
94
+
95
+ model.eval()
96
+ model.to(device=DEVICE, dtype=INFERENCE_DTYPE)
97
+ print("Model loaded successfully.")
98
+ return model
99
+
100
+ def run_inference(model, prompt_embeddings, device="cuda"):
101
+
102
+ options = DiffusionLCMGeneratorOptions(
103
+ eos_threshold=0.9,
104
+ inference_timesteps=40,
105
+ initial_noise_scale=0.6,
106
+ guidance_scale=1.5,
107
+ guidance_rescale=0.7,
108
+ epsilon_scaling=1.00045,
109
+ stop_on_repetition_cosine_threshold=0.9,
110
+ seed=42,
111
+ )
112
+
113
+ generator = TwoTowerDiffusionLCMGenerator(model, options, eos_vec=get_eos_vector())
114
+
115
+ seqs = prompt_embeddings.to(device)
116
+ batch_input = EmbeddingsBatch(seqs=seqs, padding_mask=None)
117
+
118
+ print("Running generation...")
119
+ output = generator(batch_input)
120
+
121
+ return output
122
+
123
+ if __name__ == "__main__":
124
+
125
+ raw_prompt_text = "User turn.\n\nJohn lives in his house and loves to play soccer.\n\nGive a brief definition of the word \"house\" in the sentence given as input. Generate only the definition.\n\nAssistant turn."
126
+
127
+ model = load_two_tower_model(CHECKPOINT_PATH, DEVICE)
128
+
129
+ with torch.no_grad():
130
+
131
+ sat_model = SaT("segment-any-text/sat-3l")
132
+ if torch.cuda.is_available():
133
+ sat_model.half().to(DEVICE)
134
+
135
+ split_outputs = list(sat_model.split([raw_prompt_text], threshold=0.02))
136
+ sentences = [s.strip() for s in split_outputs[0] if s.strip()]
137
+
138
+ print(sentences)
139
+
140
+ prompt = TEXT_EMBEDDER.predict(sentences, source_lang="eng_Latn", batch_size=1024)
141
+ prompt = prompt.to(device=DEVICE, dtype=INFERENCE_DTYPE)
142
+ prompt = prompt.unsqueeze(0)
143
+
144
+ results = run_inference(model, prompt, DEVICE)
145
+
146
+ for j, hyp in enumerate(results.hypotheses[0]):
147
+ print(decode_embeddings(hyp.seq)[prompt.shape[1]:])
148
+ ```
149
+
150
+ Make sure to change the source language from "eng_Latn" to the one you want to perform inference with.
151
+
152
+ Furthermore, for the instruct model, we used the following mappings for the "User turn." and "Assistant turn." strings.
153
+
154
+ ```python
155
+ USER_TRANSLATION = {
156
+ "arb_Arab": "دور المستخدم.",
157
+ "bel_Cyrl": "Ход карыстальніка.",
158
+ "ben_Beng": "ব্যবহারকারীর পালা।",
159
+ "bos_Latn": "Red korisnika.",
160
+ "bul_Cyrl": "Ред е на потребителя.",
161
+ "cat_Latn": "Torn de l'usuari.",
162
+ "ces_Latn": "Je řada na uživateli.",
163
+ "cym_Latn": "Tro'r defnyddiwr.",
164
+ "dan_Latn": "Brugerens tur.",
165
+ "deu_Latn": "Der Benutzer ist am Zug.",
166
+ "eng_Latn": "User turn.",
167
+ "fra_Latn": "C'est au tour de l'utilisateur.",
168
+ "heb_Hebr": "תור המשתמש.",
169
+ "hin_Deva": "उपयोगकर्ता की बारी।",
170
+ "hrv_Latn": "Poteg korisnika.",
171
+ "ind_Latn": "Giliran pengguna.",
172
+ "jpn_Jpan": "ユーザーの番です。",
173
+ "ita_Latn": "È il turno dell'utente.",
174
+ "kan_Knda": "ಬಳಕೆದಾರರ ತಿರುವು.",
175
+ "kor_Hang": "사용자 차례입니다.",
176
+ "lvs_Latn": "Lietotāja kārta.",
177
+ "mal_Mlym": "ഉപയോക്താവിന്റെ ഊഴം",
178
+ "mar_Deva": "वापरकर्त्याची पाळी.",
179
+ "mkd_Cyrl": "Потребен е корисник.",
180
+ "npi_Deva": "प्रयोगकर्ताको पालो।",
181
+ "nld_Latn": "De gebruiker is aan de beurt.",
182
+ "ory_Orya": "ୟୁଜର୍ ଟର୍ନ୍ |",
183
+ "pol_Latn": "Ruch użytkownika.",
184
+ "por_Latn": "É a vez do utilizador.",
185
+ "ron_Latn": "E rândul utilizatorului.",
186
+ "rus_Cyrl": "Ход пользователя.",
187
+ "slk_Latn": "Je na ťa.",
188
+ "slv_Latn": "Na vrsti je uporabnik.",
189
+ "srp_Cyrl": "Кориснички потез.",
190
+ "spa_Latn": "Turno del usuario.",
191
+ "swe_Latn": "Användarens tur.",
192
+ "swh_Latn": "Mzunguko wa mtumiaji.",
193
+ "tam_Taml": "பயனர் முறை.",
194
+ "tel_Telu": "వినియోగదారుని వంతు",
195
+ "tha_Thai": "ผู้ใช้เทิร์น",
196
+ "tur_Latn": "Kullanıcı sırası.",
197
+ "ukr_Cyrl": "Хід користувача.",
198
+ "urd_Arab": "صارف کی باری۔",
199
+ "vie_Latn": "Đến lượt người dùng.",
200
+ "zho_Hans": "轮到用户了。"
201
+ }
202
+ ```
203
+
204
+ ```python
205
+ ASSISTANT_TRANSLATION = {
206
+ "arb_Arab": "دور المساعد.",
207
+ "bel_Cyrl": "Памочнік, свой ход.",
208
+ "ben_Beng": "সহকারী পালা।",
209
+ "bos_Latn": "Pomoćni potez.",
210
+ "bul_Cyrl": "Ред е на помощника.",
211
+ "cat_Latn": "Torn de l'ajudant.",
212
+ "ces_Latn": "Na řadě je asistent.",
213
+ "cym_Latn": "Tro'r cynorthwyydd.",
214
+ "dan_Latn": "Assistenten er på tur.",
215
+ "deu_Latn": "Der Assistent ist an der Reihe.",
216
+ "eng_Latn": "Assistant turn.",
217
+ "fra_Latn": "C'est au tour de l'assistant.",
218
+ "heb_Hebr": "תורו של העוזר.",
219
+ "hin_Deva": "सहायक की बारी।",
220
+ "hrv_Latn": "Pomoćni potez.",
221
+ "ind_Latn": "Giliran asisten.",
222
+ "jpn_Jpan": "アシスタントの番。",
223
+ "ita_Latn": "È il turno dell'assistente.",
224
+ "kan_Knda": "ಸಹಾಯಕ ತಿರುವು.",
225
+ "kor_Hang": "조수 차례.",
226
+ "lvs_Latn": "Palīga kārta.",
227
+ "mal_Mlym": "സഹായിയുടെ ഊഴം",
228
+ "mar_Deva": "सहाय्यक पालवी.",
229
+ "mkd_Cyrl": "Помошник-ред.",
230
+ "npi_Deva": "सहायक पालो।",
231
+ "nld_Latn": "De assistent is aan de beurt.",
232
+ "ory_Orya": "ଆସିଷ୍ଟାଣ୍ଟ ଟର୍ନ୍ |",
233
+ "pol_Latn": "Rzuty pomocników.",
234
+ "por_Latn": "É a vez do assistente.",
235
+ "ron_Latn": "E rândul asistentului.",
236
+ "rus_Cyrl": "Очередь помощника.",
237
+ "slk_Latn": "Na rade je asistent.",
238
+ "slv_Latn": "Na vrsti je pomočnik.",
239
+ "srp_Cyrl": "Помоћни круг.",
240
+ "spa_Latn": "Turno del asistente.",
241
+ "swe_Latn": "Assistenten är på tur.",
242
+ "swh_Latn": "Mzunguko wa msaidizi.",
243
+ "tam_Taml": "துணைவரின் முறை.",
244
+ "tel_Telu": "సహాయకుడి వంతు",
245
+ "tha_Thai": "ผู้ช่วยเลี้ยว",
246
+ "tur_Latn": "Sıra asistanında.",
247
+ "ukr_Cyrl": "Черга помічника.",
248
+ "urd_Arab": "معاون کی باری۔",
249
+ "vie_Latn": "Đến lượt trợ lý.",
250
+ "zho_Hans": "助手的回合。"
251
+ }
252
+ ```
253
+
254
+ ## Citation
255
+
256
+ If you use this model in your research, please cite the following:
257
+
258
+ ```bibtex
259
+ @misc{musacchio2026mimirlargescalemultilingualconcept,
260
+ title={Mimir: Large-scale Multilingual Concept Modeling},
261
+ author={Elio Musacchio and Lucia Siciliani and Pierpaolo Basile},
262
+ year={2026},
263
+ eprint={2605.25263},
264
+ archivePrefix={arXiv},
265
+ primaryClass={cs.CL},
266
+ url={https://arxiv.org/abs/2605.25263},
267
+ }
268
+ ```