refactor: 修改模型架构
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user