chenbhao opencode commited on
Commit
91a3777
·
1 Parent(s): 45cf443

Restructure models into lm/vlm/vam sub-packages by modality

Browse files

Split the flat model modules into three modality-scoped sub-packages,
each with a config.py (configuration) and model.py (modeling code):

- models/lm/ : pure text -> MiniMindConfig + MiniMindForCausalLM
- models/vlm/ : text + vision -> VLMConfig + MiniMindVLM
- models/vam/ : text + audio -> OmniConfig + MiniMindOmni (+ TalkerModule)

models/__init__.py re-exports all symbols so existing
`from omni.models import ...` references keep working. Updated
trainers/examples/utils that imported the old flat submodule paths.
README reflects the new layout.

Co-Authored-By: opencode <noreply@opencode.ai>

README.md CHANGED
@@ -14,8 +14,10 @@ Omni 是一个以 **多模态 (omni)** 为目标的 LLM 训练 / 推理框架,
14
  `norm.py`(`RMSNorm`)、`rope.py`(`precompute_freqs_cis` / `apply_rotary_pos_emb` / `repeat_kv`)、
15
  `attention.py`(`Attention`)、`mlp.py`(`FeedForward` / `MOEFeedForward`)、
16
  `block.py`(`MiniMindBlock`)、`model.py`(`MiniMindModel` Transformer 主干)。
17
- - `models/`:把 `core` 组件**拼装**成成品模型——
18
- `MiniMindForCausalLM`文本)、`MiniMindVLM`(视觉)、`MiniMindOmni`(语音/全模态)。
 
 
19
  - `encoders/`:外部模态编码器,按模态分目录——`vision/`(SigLIP)、`audio/`(SenseVoice)。
20
  - `projectors/`:把 encoder 输出**桥接**到 LLM 隐藏维度的拼接层(`MMVisionProjector`、`MMAudioProjector`)。
21
  - `serve/`:实时语音会话工程层(`SileroVAD`、`RealtimeSession`)。
@@ -31,10 +33,16 @@ src/omni/
31
  │ ├── mlp.py # FeedForward / MOEFeedForward
32
  │ ├── block.py # MiniMindBlock
33
  │ └── model.py # MiniMindModel(Transformer 主干)
34
- ├── models/ # 模型拼装
35
- │ ├── minimind.py # MiniMindConfig + MiniMindForCausalLM
36
- │ ├── vlm.py # VLMConfig + MiniMindVLM(视觉多模态)
37
- ── omni.py # OmniConfig + MiniMindOmni + TalkerModule(语音/全模态)
 
 
 
 
 
 
38
  │ └── lora.py # LoRA 注入 / 保存 / 合并
39
  ├── encoders/ # 多模态编码器(按模态分目录)
40
  │ ├── vision/ # SiglipVisionEncoder
 
14
  `norm.py`(`RMSNorm`)、`rope.py`(`precompute_freqs_cis` / `apply_rotary_pos_emb` / `repeat_kv`)、
15
  `attention.py`(`Attention`)、`mlp.py`(`FeedForward` / `MOEFeedForward`)、
16
  `block.py`(`MiniMindBlock`)、`model.py`(`MiniMindModel` Transformer 主干)。
17
+ - `models/`:把 `core` 组件**拼装**成成品模型,按模态能力分为三个子包(每个含 `config.py` 配置 + `model.py` 建模):
18
+ - `models/lm/`:纯文本——`MiniMindConfig` + `MiniMindForCausalLM`
19
+ - `models/vlm/`:文本 + 视觉——`VLMConfig` + `MiniMindVLM`
20
+ - `models/vam/`:文本 + 语音/全模态——`OmniConfig` + `MiniMindOmni`(含 `TalkerModule`)
21
  - `encoders/`:外部模态编码器,按模态分目录——`vision/`(SigLIP)、`audio/`(SenseVoice)。
22
  - `projectors/`:把 encoder 输出**桥接**到 LLM 隐藏维度的拼接层(`MMVisionProjector`、`MMAudioProjector`)。
23
  - `serve/`:实时语音会话工程层(`SileroVAD`、`RealtimeSession`)。
 
33
  │ ├── mlp.py # FeedForward / MOEFeedForward
34
  │ ├── block.py # MiniMindBlock
35
  │ └── model.py # MiniMindModel(Transformer 主干)
36
+ ├── models/ # 模型拼装(按模态能力分子包)
37
+ │ ├── lm/ # 纯文本
38
+ ├── config.py # MiniMindConfig
39
+ │ └── model.py # MiniMindForCausalLM
40
+ │ ├── vlm/ # 文本 + 视觉
41
+ │ │ ├── config.py # VLMConfig
42
+ │ │ └── model.py # MiniMindVLM
43
+ │ ├── vam/ # 文本 + 语音/全模态
44
+ │ │ ├── config.py # OmniConfig
45
+ │ │ └── model.py # MiniMindOmni + TalkerModule
46
  │ └── lora.py # LoRA 注入 / 保存 / 合并
47
  ├── encoders/ # 多模态编码器(按模态分目录)
48
  │ ├── vision/ # SiglipVisionEncoder
examples/convert_model.py CHANGED
@@ -5,7 +5,7 @@ import torch
5
  import transformers
6
  import warnings
7
  from transformers import AutoTokenizer, AutoModelForCausalLM, Qwen3Config, Qwen3ForCausalLM, Qwen3MoeConfig, Qwen3MoeForCausalLM
8
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
9
  from omni.models.lora import apply_lora, merge_lora
10
 
11
  warnings.filterwarnings('ignore', category=UserWarning)
 
5
  import transformers
6
  import warnings
7
  from transformers import AutoTokenizer, AutoModelForCausalLM, Qwen3Config, Qwen3ForCausalLM, Qwen3MoeConfig, Qwen3MoeForCausalLM
8
+ from omni.models import MiniMindConfig, MiniMindForCausalLM
9
  from omni.models.lora import apply_lora, merge_lora
10
 
11
  warnings.filterwarnings('ignore', category=UserWarning)
examples/eval_llm.py CHANGED
@@ -4,7 +4,7 @@ import random
4
  import warnings
5
  import torch
6
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
7
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
8
  from omni.models.lora import * # noqa: F401,F403
9
  from omni.utils.training import setup_seed, get_model_params
10
  warnings.filterwarnings('ignore')
 
4
  import warnings
5
  import torch
6
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
7
+ from omni.models import MiniMindConfig, MiniMindForCausalLM
8
  from omni.models.lora import * # noqa: F401,F403
9
  from omni.utils.training import setup_seed, get_model_params
10
  warnings.filterwarnings('ignore')
examples/eval_toolcall.py CHANGED
@@ -9,7 +9,7 @@ import torch
9
  from datetime import datetime
10
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
11
  from openai import OpenAI
12
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
13
  from omni.utils.training import setup_seed, get_model_params
14
  warnings.filterwarnings('ignore')
15
 
 
9
  from datetime import datetime
10
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
11
  from openai import OpenAI
12
+ from omni.models import MiniMindConfig, MiniMindForCausalLM
13
  from omni.utils.training import setup_seed, get_model_params
14
  warnings.filterwarnings('ignore')
15
 
examples/serve_openai_api.py CHANGED
@@ -14,7 +14,7 @@ from fastapi import FastAPI, HTTPException
14
  from fastapi.responses import StreamingResponse
15
  from pydantic import BaseModel, Field
16
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
17
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
18
  from omni.models.lora import apply_lora, load_lora
19
 
20
  warnings.filterwarnings('ignore')
 
14
  from fastapi.responses import StreamingResponse
15
  from pydantic import BaseModel, Field
16
  from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer
17
+ from omni.models import MiniMindConfig, MiniMindForCausalLM
18
  from omni.models.lora import apply_lora, load_lora
19
 
20
  warnings.filterwarnings('ignore')
src/omni/models/__init__.py CHANGED
@@ -1,3 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  from omni.core import (
2
  RMSNorm,
3
  Attention,
@@ -8,46 +21,26 @@ from omni.core import (
8
  precompute_freqs_cis,
9
  apply_rotary_pos_emb,
10
  )
11
- from omni.models.minimind import (
12
- MiniMindConfig,
13
- MiniMindForCausalLM,
14
- )
15
- from omni.models.vlm import (
16
- VLMConfig,
17
- MiniMindVLM,
18
- )
19
- from omni.models.omni import (
20
- OmniConfig,
21
- MiniMindOmni,
22
- TalkerModule,
23
- )
24
- from omni.models.lora import (
25
- LoRA,
26
- apply_lora,
27
- load_lora,
28
- save_lora,
29
- merge_lora,
30
- )
31
 
32
  __all__ = [
33
  "MiniMindConfig",
34
- "MiniMindModel",
35
  "MiniMindForCausalLM",
36
  "VLMConfig",
37
  "MiniMindVLM",
38
  "OmniConfig",
39
  "MiniMindOmni",
40
  "TalkerModule",
 
 
 
 
 
41
  "RMSNorm",
42
  "Attention",
43
  "FeedForward",
44
  "MOEFeedForward",
45
  "MiniMindBlock",
 
46
  "precompute_freqs_cis",
47
  "apply_rotary_pos_emb",
48
- "LoRA",
49
- "apply_lora",
50
- "load_lora",
51
- "save_lora",
52
- "merge_lora",
53
  ]
 
1
+ from omni.models.lm.config import MiniMindConfig
2
+ from omni.models.lm.model import MiniMindForCausalLM
3
+ from omni.models.vlm.config import VLMConfig
4
+ from omni.models.vlm.model import MiniMindVLM
5
+ from omni.models.vam.config import OmniConfig
6
+ from omni.models.vam.model import MiniMindOmni, TalkerModule
7
+ from omni.models.lora import (
8
+ LoRA,
9
+ apply_lora,
10
+ load_lora,
11
+ save_lora,
12
+ merge_lora,
13
+ )
14
  from omni.core import (
15
  RMSNorm,
16
  Attention,
 
21
  precompute_freqs_cis,
22
  apply_rotary_pos_emb,
23
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
 
25
  __all__ = [
26
  "MiniMindConfig",
 
27
  "MiniMindForCausalLM",
28
  "VLMConfig",
29
  "MiniMindVLM",
30
  "OmniConfig",
31
  "MiniMindOmni",
32
  "TalkerModule",
33
+ "LoRA",
34
+ "apply_lora",
35
+ "load_lora",
36
+ "save_lora",
37
+ "merge_lora",
38
  "RMSNorm",
39
  "Attention",
40
  "FeedForward",
41
  "MOEFeedForward",
42
  "MiniMindBlock",
43
+ "MiniMindModel",
44
  "precompute_freqs_cis",
45
  "apply_rotary_pos_emb",
 
 
 
 
 
46
  ]
src/omni/models/lm/__init__.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ from omni.models.lm.config import MiniMindConfig
2
+ from omni.models.lm.model import MiniMindForCausalLM
3
+
4
+ __all__ = ["MiniMindConfig", "MiniMindForCausalLM"]
src/omni/models/lm/config.py ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from transformers import PretrainedConfig
3
+
4
+
5
+ class MiniMindConfig(PretrainedConfig):
6
+ model_type = "minimind"
7
+
8
+ def __init__(self, hidden_size=768, num_hidden_layers=8, use_moe=False, **kwargs):
9
+ super().__init__(**kwargs)
10
+ self.hidden_size = hidden_size
11
+ self.num_hidden_layers = num_hidden_layers
12
+ self.use_moe = use_moe
13
+ self.dropout = kwargs.get("dropout", 0.0)
14
+ self.vocab_size = kwargs.get("vocab_size", 6400)
15
+ self.bos_token_id = kwargs.get("bos_token_id", 1)
16
+ self.eos_token_id = kwargs.get("eos_token_id", 2)
17
+ self.flash_attn = kwargs.get("flash_attn", True)
18
+ self.num_attention_heads = kwargs.get("num_attention_heads", 8)
19
+ self.num_key_value_heads = kwargs.get("num_key_value_heads", 4)
20
+ self.head_dim = kwargs.get("head_dim", self.hidden_size // self.num_attention_heads)
21
+ self.hidden_act = kwargs.get("hidden_act", 'silu')
22
+ self.intermediate_size = kwargs.get("intermediate_size", math.ceil(hidden_size * math.pi / 64) * 64)
23
+ self.max_position_embeddings = kwargs.get("max_position_embeddings", 32768)
24
+ self.rms_norm_eps = kwargs.get("rms_norm_eps", 1e-6)
25
+ self.rope_theta = kwargs.get("rope_theta", 1e6)
26
+ self.tie_word_embeddings = kwargs.get("tie_word_embeddings", True)
27
+ self.inference_rope_scaling = kwargs.get("inference_rope_scaling", False)
28
+ self.rope_scaling = {
29
+ "beta_fast": 32,
30
+ "beta_slow": 1,
31
+ "factor": 16,
32
+ "original_max_position_embeddings": 2048,
33
+ "attention_factor": 1.0,
34
+ "type": "yarn"
35
+ } if self.inference_rope_scaling else None
36
+ # MoE specific configs (ignored if use_moe = False)
37
+ self.num_experts = kwargs.get("num_experts", 4)
38
+ self.num_experts_per_tok = kwargs.get("num_experts_per_tok", 1)
39
+ self.moe_intermediate_size = kwargs.get("moe_intermediate_size", self.intermediate_size)
40
+ self.norm_topk_prob = kwargs.get("norm_topk_prob", True)
41
+ self.router_aux_loss_coef = kwargs.get("router_aux_loss_coef", 5e-4)
src/omni/models/{minimind.py → lm/model.py} RENAMED
@@ -1,50 +1,11 @@
1
- import math
2
  import torch
3
  import torch.nn.functional as F
4
  from torch import nn
5
- from transformers import PreTrainedModel, GenerationMixin, PretrainedConfig
6
  from transformers.modeling_outputs import MoeCausalLMOutputWithPast
7
 
8
  from omni.core import MiniMindModel
9
-
10
-
11
- class MiniMindConfig(PretrainedConfig):
12
- model_type = "minimind"
13
-
14
- def __init__(self, hidden_size=768, num_hidden_layers=8, use_moe=False, **kwargs):
15
- super().__init__(**kwargs)
16
- self.hidden_size = hidden_size
17
- self.num_hidden_layers = num_hidden_layers
18
- self.use_moe = use_moe
19
- self.dropout = kwargs.get("dropout", 0.0)
20
- self.vocab_size = kwargs.get("vocab_size", 6400)
21
- self.bos_token_id = kwargs.get("bos_token_id", 1)
22
- self.eos_token_id = kwargs.get("eos_token_id", 2)
23
- self.flash_attn = kwargs.get("flash_attn", True)
24
- self.num_attention_heads = kwargs.get("num_attention_heads", 8)
25
- self.num_key_value_heads = kwargs.get("num_key_value_heads", 4)
26
- self.head_dim = kwargs.get("head_dim", self.hidden_size // self.num_attention_heads)
27
- self.hidden_act = kwargs.get("hidden_act", 'silu')
28
- self.intermediate_size = kwargs.get("intermediate_size", math.ceil(hidden_size * math.pi / 64) * 64)
29
- self.max_position_embeddings = kwargs.get("max_position_embeddings", 32768)
30
- self.rms_norm_eps = kwargs.get("rms_norm_eps", 1e-6)
31
- self.rope_theta = kwargs.get("rope_theta", 1e6)
32
- self.tie_word_embeddings = kwargs.get("tie_word_embeddings", True)
33
- self.inference_rope_scaling = kwargs.get("inference_rope_scaling", False)
34
- self.rope_scaling = {
35
- "beta_fast": 32,
36
- "beta_slow": 1,
37
- "factor": 16,
38
- "original_max_position_embeddings": 2048,
39
- "attention_factor": 1.0,
40
- "type": "yarn"
41
- } if self.inference_rope_scaling else None
42
- # MoE specific configs (ignored if use_moe = False)
43
- self.num_experts = kwargs.get("num_experts", 4)
44
- self.num_experts_per_tok = kwargs.get("num_experts_per_tok", 1)
45
- self.moe_intermediate_size = kwargs.get("moe_intermediate_size", self.intermediate_size)
46
- self.norm_topk_prob = kwargs.get("norm_topk_prob", True)
47
- self.router_aux_loss_coef = kwargs.get("router_aux_loss_coef", 5e-4)
48
 
49
 
50
  class MiniMindForCausalLM(PreTrainedModel, GenerationMixin):
 
 
1
  import torch
2
  import torch.nn.functional as F
3
  from torch import nn
4
+ from transformers import PreTrainedModel, GenerationMixin
5
  from transformers.modeling_outputs import MoeCausalLMOutputWithPast
6
 
7
  from omni.core import MiniMindModel
8
+ from omni.models.lm.config import MiniMindConfig
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
 
10
 
11
  class MiniMindForCausalLM(PreTrainedModel, GenerationMixin):
src/omni/models/vam/__init__.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ from omni.models.vam.config import OmniConfig
2
+ from omni.models.vam.model import MiniMindOmni, TalkerModule
3
+
4
+ __all__ = ["OmniConfig", "MiniMindOmni", "TalkerModule"]
src/omni/models/vam/config.py ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from omni.models.lm.config import MiniMindConfig
2
+
3
+
4
+ class OmniConfig(MiniMindConfig):
5
+ model_type = "minimind-o"
6
+
7
+ def __init__(self, **kwargs):
8
+ super().__init__(**kwargs)
9
+ self.num_talker_hidden_layers = kwargs.get("num_talker_hidden_layers", 4)
10
+ self.talker_hidden_size = kwargs.get("talker_hidden_size", 768)
11
+ self.audio_ids = kwargs.get("audio_ids", [16])
12
+ self.audio_special_token = kwargs.get("audio_special_token", "<|audio_pad|>")
13
+ self.audio_hidden_size = kwargs.get("audio_hidden_size", 512)
14
+ self.audio_vocab_size = kwargs.get("audio_vocab_size", 2112)
15
+ self.audio_pad_token = kwargs.get("audio_pad_token", 2049)
16
+ self.audio_stop_token = kwargs.get("audio_stop_token", 2050)
17
+ self.audio_spk_token = kwargs.get("audio_spk_token", 2051)
18
+ self.spk_emb_size = kwargs.get("spk_emb_size", 192)
19
+ self.think_end_ids = kwargs.get("think_end_ids", [26, 234, 234])
20
+ self.image_ids = kwargs.get("image_ids", [12])
21
+ self.image_special_token = kwargs.get("image_special_token", "<|image_pad|>")
22
+ self.image_hidden_size = kwargs.get("image_hidden_size", 768)
23
+ self.image_token_len = kwargs.get("image_token_len", 64)
24
+ self.bridge_layer = kwargs.get("bridge_layer", self.num_hidden_layers // 2 - 1)
src/omni/models/{omni.py → vam/model.py} RENAMED
@@ -10,36 +10,15 @@ from torch.nn import functional as F
10
  from transformers.modeling_outputs import MoeCausalLMOutputWithPast
11
  from transformers import SiglipVisionModel, SiglipImageProcessor, logging as hf_logging
12
 
13
- from omni.core import RMSNorm, precompute_freqs_cis, apply_rotary_pos_emb, MiniMindBlock, MOEFeedForward
14
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
15
- from omni.encoders.audio import SenseVoiceAudioEncoder
 
 
16
  from omni.encoders.vision import SiglipVisionEncoder
17
  from omni.projectors import MMVisionProjector, MMAudioProjector
18
 
19
 
20
- class OmniConfig(MiniMindConfig):
21
- model_type = "minimind-o"
22
-
23
- def __init__(self, **kwargs):
24
- super().__init__(**kwargs)
25
- self.num_talker_hidden_layers = kwargs.get("num_talker_hidden_layers", 4)
26
- self.talker_hidden_size = kwargs.get("talker_hidden_size", 768)
27
- self.audio_ids = kwargs.get("audio_ids", [16])
28
- self.audio_special_token = kwargs.get("audio_special_token", "<|audio_pad|>")
29
- self.audio_hidden_size = kwargs.get("audio_hidden_size", 512)
30
- self.audio_vocab_size = kwargs.get("audio_vocab_size", 2112)
31
- self.audio_pad_token = kwargs.get("audio_pad_token", 2049)
32
- self.audio_stop_token = kwargs.get("audio_stop_token", 2050)
33
- self.audio_spk_token = kwargs.get("audio_spk_token", 2051)
34
- self.spk_emb_size = kwargs.get("spk_emb_size", 192)
35
- self.think_end_ids = kwargs.get("think_end_ids", [26, 234, 234])
36
- self.image_ids = kwargs.get("image_ids", [12])
37
- self.image_special_token = kwargs.get("image_special_token", "<|image_pad|>")
38
- self.image_hidden_size = kwargs.get("image_hidden_size", 768)
39
- self.image_token_len = kwargs.get("image_token_len", 64)
40
- self.bridge_layer = kwargs.get("bridge_layer", self.num_hidden_layers // 2 - 1)
41
-
42
-
43
  class TalkerHead(nn.Module):
44
  def __init__(self, in_features, out_features, num_layers=8, rank=256):
45
  super().__init__()
 
10
  from transformers.modeling_outputs import MoeCausalLMOutputWithPast
11
  from transformers import SiglipVisionModel, SiglipImageProcessor, logging as hf_logging
12
 
13
+ from omni.core import RMSNorm, precompute_freqs_cis, MiniMindBlock, MOEFeedForward
14
+ from omni.models.lm.config import MiniMindConfig
15
+ from omni.models.lm.model import MiniMindForCausalLM
16
+ from omni.models.vam.config import OmniConfig
17
+ from omni.encoders.audio import SenseVoiceAudioEncoder, SenseVoiceAudioProcessor
18
  from omni.encoders.vision import SiglipVisionEncoder
19
  from omni.projectors import MMVisionProjector, MMAudioProjector
20
 
21
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
22
  class TalkerHead(nn.Module):
23
  def __init__(self, in_features, out_features, num_layers=8, rank=256):
24
  super().__init__()
src/omni/models/vlm/__init__.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ from omni.models.vlm.config import VLMConfig
2
+ from omni.models.vlm.model import MiniMindVLM
3
+
4
+ __all__ = ["VLMConfig", "MiniMindVLM"]
src/omni/models/vlm/config.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from omni.models.lm.config import MiniMindConfig
2
+
3
+
4
+ class VLMConfig(MiniMindConfig):
5
+ model_type = "minimind-v"
6
+
7
+ def __init__(self, image_special_token='<|image_pad|>', image_ids=[12], **kwargs):
8
+ self.image_special_token = image_special_token
9
+ self.image_ids = image_ids
10
+ self.image_hidden_size = kwargs.get("image_hidden_size", 768)
11
+ self.image_token_len = kwargs.get("image_token_len", 64)
12
+ super().__init__(**kwargs)
src/omni/models/{vlm.py → vlm/model.py} RENAMED
@@ -1,4 +1,3 @@
1
- import os
2
  import torch
3
  import torch.nn.functional as F
4
  import warnings
@@ -7,24 +6,14 @@ from torch import nn
7
  from transformers.modeling_outputs import MoeCausalLMOutputWithPast
8
 
9
  from omni.core import precompute_freqs_cis, MOEFeedForward
10
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
 
11
  from omni.encoders.vision import SiglipVisionEncoder
12
  from omni.projectors import MMVisionProjector
13
 
14
  warnings.filterwarnings('ignore')
15
 
16
 
17
- class VLMConfig(MiniMindConfig):
18
- model_type = "minimind-v"
19
-
20
- def __init__(self, image_special_token='<|image_pad|>', image_ids=[12], **kwargs):
21
- self.image_special_token = image_special_token
22
- self.image_ids = image_ids
23
- self.image_hidden_size = kwargs.get("image_hidden_size", 768)
24
- self.image_token_len = kwargs.get("image_token_len", 64)
25
- super().__init__(**kwargs)
26
-
27
-
28
  class MiniMindVLM(MiniMindForCausalLM):
29
  config_class = VLMConfig
30
 
 
 
1
  import torch
2
  import torch.nn.functional as F
3
  import warnings
 
6
  from transformers.modeling_outputs import MoeCausalLMOutputWithPast
7
 
8
  from omni.core import precompute_freqs_cis, MOEFeedForward
9
+ from omni.models.lm.model import MiniMindForCausalLM
10
+ from omni.models.vlm.config import VLMConfig
11
  from omni.encoders.vision import SiglipVisionEncoder
12
  from omni.projectors import MMVisionProjector
13
 
14
  warnings.filterwarnings('ignore')
15
 
16
 
 
 
 
 
 
 
 
 
 
 
 
17
  class MiniMindVLM(MiniMindForCausalLM):
18
  config_class = VLMConfig
19
 
src/omni/trainers/agent.py CHANGED
@@ -19,7 +19,7 @@ from torch.nn.parallel import DistributedDataParallel
19
  from torch.utils.data import DataLoader, DistributedSampler
20
  from torch.optim.lr_scheduler import CosineAnnealingLR
21
  from transformers import AutoTokenizer
22
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
23
  from omni.datasets.lm_dataset import AgentRLDataset
24
  from omni.utils.training import Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, SkipBatchSampler, init_model, LMForRewardModel
25
  from omni.trainers.rollout_engine import create_rollout_engine, compute_per_token_logps
 
19
  from torch.utils.data import DataLoader, DistributedSampler
20
  from torch.optim.lr_scheduler import CosineAnnealingLR
21
  from transformers import AutoTokenizer
22
+ from omni.models import MiniMindConfig, MiniMindForCausalLM
23
  from omni.datasets.lm_dataset import AgentRLDataset
24
  from omni.utils.training import Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, SkipBatchSampler, init_model, LMForRewardModel
25
  from omni.trainers.rollout_engine import create_rollout_engine, compute_per_token_logps
src/omni/trainers/distillation.py CHANGED
@@ -12,7 +12,7 @@ from contextlib import nullcontext
12
  from torch import optim
13
  from torch.nn.parallel import DistributedDataParallel
14
  from torch.utils.data import DataLoader, DistributedSampler
15
- from omni.models.minimind import MiniMindConfig
16
  from omni.datasets.lm_dataset import SFTDataset
17
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
18
 
 
12
  from torch import optim
13
  from torch.nn.parallel import DistributedDataParallel
14
  from torch.utils.data import DataLoader, DistributedSampler
15
+ from omni.models import MiniMindConfig
16
  from omni.datasets.lm_dataset import SFTDataset
17
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
18
 
src/omni/trainers/dpo.py CHANGED
@@ -12,7 +12,7 @@ from contextlib import nullcontext
12
  from torch import optim
13
  from torch.nn.parallel import DistributedDataParallel
14
  from torch.utils.data import DataLoader, DistributedSampler
15
- from omni.models.minimind import MiniMindConfig
16
  from omni.datasets.lm_dataset import DPODataset
17
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
18
 
 
12
  from torch import optim
13
  from torch.nn.parallel import DistributedDataParallel
14
  from torch.utils.data import DataLoader, DistributedSampler
15
+ from omni.models import MiniMindConfig
16
  from omni.datasets.lm_dataset import DPODataset
17
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
18
 
src/omni/trainers/full_sft.py CHANGED
@@ -11,7 +11,7 @@ from contextlib import nullcontext
11
  from torch import optim, nn
12
  from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
- from omni.models.minimind import MiniMindConfig
15
  from omni.datasets.lm_dataset import SFTDataset
16
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
17
 
 
11
  from torch import optim, nn
12
  from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
+ from omni.models import MiniMindConfig
15
  from omni.datasets.lm_dataset import SFTDataset
16
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
17
 
src/omni/trainers/grpo.py CHANGED
@@ -17,7 +17,7 @@ from torch.nn.parallel import DistributedDataParallel
17
  from torch.utils.data import DataLoader, DistributedSampler
18
  from torch.optim.lr_scheduler import CosineAnnealingLR
19
  from transformers import AutoModel
20
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
21
  from omni.datasets.lm_dataset import RLAIFDataset
22
  from omni.utils.training import Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, SkipBatchSampler, init_model, LMForRewardModel
23
  from omni.trainers.rollout_engine import create_rollout_engine
 
17
  from torch.utils.data import DataLoader, DistributedSampler
18
  from torch.optim.lr_scheduler import CosineAnnealingLR
19
  from transformers import AutoModel
20
+ from omni.models import MiniMindConfig, MiniMindForCausalLM
21
  from omni.datasets.lm_dataset import RLAIFDataset
22
  from omni.utils.training import Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, SkipBatchSampler, init_model, LMForRewardModel
23
  from omni.trainers.rollout_engine import create_rollout_engine
src/omni/trainers/lora.py CHANGED
@@ -11,7 +11,7 @@ from contextlib import nullcontext
11
  from torch import optim, nn
12
  from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
- from omni.models.minimind import MiniMindConfig
15
  from omni.datasets.lm_dataset import SFTDataset
16
  from omni.models.lora import save_lora, apply_lora
17
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
 
11
  from torch import optim, nn
12
  from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
+ from omni.models import MiniMindConfig
15
  from omni.datasets.lm_dataset import SFTDataset
16
  from omni.models.lora import save_lora, apply_lora
17
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
src/omni/trainers/ppo.py CHANGED
@@ -16,7 +16,7 @@ from torch.nn.parallel import DistributedDataParallel
16
  from torch.utils.data import DataLoader, DistributedSampler
17
  from torch.nn.utils import clip_grad_norm_
18
  from torch.optim.lr_scheduler import CosineAnnealingLR
19
- from omni.models.minimind import MiniMindConfig, MiniMindForCausalLM
20
  from omni.datasets.lm_dataset import RLAIFDataset
21
  from omni.utils.training import Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, SkipBatchSampler, init_model, LMForRewardModel
22
  from omni.trainers.rollout_engine import create_rollout_engine
 
16
  from torch.utils.data import DataLoader, DistributedSampler
17
  from torch.nn.utils import clip_grad_norm_
18
  from torch.optim.lr_scheduler import CosineAnnealingLR
19
+ from omni.models import MiniMindConfig, MiniMindForCausalLM
20
  from omni.datasets.lm_dataset import RLAIFDataset
21
  from omni.utils.training import Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, SkipBatchSampler, init_model, LMForRewardModel
22
  from omni.trainers.rollout_engine import create_rollout_engine
src/omni/trainers/pretrain.py CHANGED
@@ -11,7 +11,7 @@ from contextlib import nullcontext
11
  from torch import optim, nn
12
  from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
- from omni.models.minimind import MiniMindConfig
15
  from omni.datasets.lm_dataset import PretrainDataset
16
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
17
 
 
11
  from torch import optim, nn
12
  from torch.nn.parallel import DistributedDataParallel
13
  from torch.utils.data import DataLoader, DistributedSampler
14
+ from omni.models import MiniMindConfig
15
  from omni.datasets.lm_dataset import PretrainDataset
16
  from omni.utils.training import get_lr, Logger, is_main_process, lm_checkpoint, init_distributed_mode, setup_seed, init_model, SkipBatchSampler
17
 
src/omni/utils/multimodal.py CHANGED
@@ -5,7 +5,7 @@ import torch.distributed as dist
5
  from torch.nn.parallel import DistributedDataParallel
6
 
7
  from omni.utils.training import Logger, is_main_process
8
- from omni.models import MiniMindOmni
9
 
10
 
11
  def get_vlm_model_params(model, config, ignore_patterns=('vision_encoder',)):
@@ -28,7 +28,6 @@ def get_vlm_model_params(model, config, ignore_patterns=('vision_encoder',)):
28
 
29
  def init_vlm_model(vlm_config, from_weight='pretrain_vlm', tokenizer_path='../model', vision_model_path='../model/siglip2-base-p32-256-ve', save_dir='../out', device='cuda', freeze_llm=0):
30
  from transformers import AutoTokenizer
31
- from omni.models.vlm import MiniMindVLM
32
  tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
33
  model = MiniMindVLM(vlm_config, vision_model_path=vision_model_path)
34
 
 
5
  from torch.nn.parallel import DistributedDataParallel
6
 
7
  from omni.utils.training import Logger, is_main_process
8
+ from omni.models import MiniMindOmni, MiniMindVLM
9
 
10
 
11
  def get_vlm_model_params(model, config, ignore_patterns=('vision_encoder',)):
 
28
 
29
  def init_vlm_model(vlm_config, from_weight='pretrain_vlm', tokenizer_path='../model', vision_model_path='../model/siglip2-base-p32-256-ve', save_dir='../out', device='cuda', freeze_llm=0):
30
  from transformers import AutoTokenizer
 
31
  tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
32
  model = MiniMindVLM(vlm_config, vision_model_path=vision_model_path)
33
 
src/omni/utils/training.py CHANGED
@@ -8,7 +8,7 @@ import torch.distributed as dist
8
  from torch.nn.parallel import DistributedDataParallel
9
  from torch.utils.data import Sampler
10
  from transformers import AutoTokenizer, AutoModel, AutoModelForSequenceClassification
11
- from omni.models.minimind import MiniMindForCausalLM
12
 
13
 
14
  def get_model_params(model, config):
 
8
  from torch.nn.parallel import DistributedDataParallel
9
  from torch.utils.data import Sampler
10
  from transformers import AutoTokenizer, AutoModel, AutoModelForSequenceClassification
11
+ from omni.models import MiniMindForCausalLM
12
 
13
 
14
  def get_model_params(model, config):