refactor: 修改模型架构

This commit is contained in:
2026-04-16 18:57:21 +08:00
parent cb93f8219e
commit a38334f4ce
17 changed files with 540 additions and 176 deletions
+2 -2
View File
@@ -5,7 +5,7 @@ from typing import Dict, List, Any, Optional
import torch
from torch import Tensor
from pipeline.tokenizer import BpeTokenizer
from pipeline.tokenize import AutoTokenizer
from pipeline.strategies import PromptStrategy, ChatMLStrategy
from pipeline.processors.base import BaseProcessor, _encode_with_mask
from pipeline.processors.factory import ProcessorFactory
@@ -20,7 +20,7 @@ class DPOProcessor(BaseProcessor):
def __init__(
self,
tokenizer: BpeTokenizer,
tokenizer: AutoTokenizer,
strategy: Optional[PromptStrategy] = None,
):
self.tokenizer = tokenizer
+4 -4
View File
@@ -3,7 +3,7 @@
from typing import Dict, List, Any, Optional, Type
from pipeline.processors.base import BaseProcessor
from pipeline.tokenizer import BpeTokenizer
from pipeline.tokenize import AutoTokenizer
from pipeline.strategies import PromptStrategy, StrategyFactory
@@ -45,7 +45,7 @@ class ProcessorFactory:
return decorator
@classmethod
def create(cls, processor_type: str, tokenizer: BpeTokenizer) -> BaseProcessor:
def create(cls, processor_type: str, tokenizer: AutoTokenizer) -> BaseProcessor:
"""Create a processor by type name (uses default ChatMLStrategy for SFT/DPO).
Args:
@@ -69,7 +69,7 @@ class ProcessorFactory:
def create_with_strategy(
cls,
processor_type: str,
tokenizer: BpeTokenizer,
tokenizer: AutoTokenizer,
strategy: PromptStrategy,
) -> BaseProcessor:
"""Create a processor with a custom strategy.
@@ -99,7 +99,7 @@ class ProcessorFactory:
def create_with_strategy_name(
cls,
processor_type: str,
tokenizer: BpeTokenizer,
tokenizer: AutoTokenizer,
strategy_name: str,
**strategy_kwargs,
) -> BaseProcessor:
+2 -2
View File
@@ -5,7 +5,7 @@ from typing import Dict, List, Any
import torch
from torch import Tensor
from pipeline.tokenizer import BpeTokenizer
from pipeline.tokenize import AutoTokenizer
from pipeline.processors.base import BaseProcessor
from pipeline.processors.factory import ProcessorFactory
@@ -14,7 +14,7 @@ from pipeline.processors.factory import ProcessorFactory
class PreTrainProcessor(BaseProcessor):
"""Pre-training data processor."""
def __init__(self, tokenizer: BpeTokenizer):
def __init__(self, tokenizer: AutoTokenizer):
self.tokenizer = tokenizer
def process(self, input_dict: Dict[str, Any]) -> Dict[str, Tensor]:
+2 -2
View File
@@ -5,7 +5,7 @@ from typing import Dict, List, Any, Optional
import torch
from torch import Tensor
from pipeline.tokenizer import BpeTokenizer
from pipeline.tokenize import AutoTokenizer
from pipeline.strategies import PromptStrategy, ChatMLStrategy
from pipeline.processors.base import BaseProcessor, _encode_with_mask
from pipeline.processors.factory import ProcessorFactory
@@ -20,7 +20,7 @@ class SFTProcessor(BaseProcessor):
def __init__(
self,
tokenizer: BpeTokenizer,
tokenizer: AutoTokenizer,
strategy: Optional[PromptStrategy] = None,
):
self.tokenizer = tokenizer