This commit is contained in:
@@ -0,0 +1,41 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from swift.utils.import_utils import _LazyModule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
# Recommend using `xxx_main`
|
||||
from .app import app_main
|
||||
from .base import SwiftPipeline
|
||||
from .eval import eval_main
|
||||
from .export import export_main, export_to_ollama, merge_lora, quantize_model
|
||||
from .infer import deploy_main, infer_main, rollout_main, run_deploy
|
||||
from .sampling import sampling_main
|
||||
from .train import SwiftSft, pretrain_main, rlhf_main, sft_main
|
||||
from .utils import prepare_model_template
|
||||
else:
|
||||
_import_structure = {
|
||||
'infer': [
|
||||
'deploy_main',
|
||||
'infer_main',
|
||||
'run_deploy',
|
||||
'rollout_main',
|
||||
],
|
||||
'export': ['export_main', 'merge_lora', 'quantize_model', 'export_to_ollama'],
|
||||
'app': ['app_main'],
|
||||
'eval': ['eval_main'],
|
||||
'train': ['sft_main', 'pretrain_main', 'rlhf_main', 'SwiftSft'],
|
||||
'sampling': ['sampling_main'],
|
||||
'base': ['SwiftPipeline'],
|
||||
'utils': ['prepare_model_template'],
|
||||
}
|
||||
|
||||
import sys
|
||||
|
||||
sys.modules[__name__] = _LazyModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
from .app import SwiftApp, app_main
|
||||
@@ -0,0 +1,43 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import gradio
|
||||
from contextlib import nullcontext
|
||||
from packaging import version
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import AppArguments
|
||||
from swift.utils import get_logger
|
||||
from ..base import SwiftPipeline
|
||||
from ..infer import run_deploy
|
||||
from .build_ui import build_ui
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftApp(SwiftPipeline):
|
||||
args_class = AppArguments
|
||||
args: args_class
|
||||
|
||||
def run(self):
|
||||
args = self.args
|
||||
deploy_context = nullcontext() if args.base_url else run_deploy(args, return_url=True)
|
||||
with deploy_context as base_url:
|
||||
base_url = base_url or args.base_url
|
||||
demo = build_ui(
|
||||
base_url,
|
||||
args.model_suffix,
|
||||
request_config=args.get_request_config(),
|
||||
is_multimodal=args.is_multimodal,
|
||||
studio_title=args.studio_title,
|
||||
lang=args.lang,
|
||||
default_system=args.system)
|
||||
concurrency_count = 1 if args.infer_backend == 'transformers' else 16
|
||||
if version.parse(gradio.__version__) < version.parse('4'):
|
||||
queue_kwargs = {'concurrency_count': concurrency_count}
|
||||
else:
|
||||
queue_kwargs = {'default_concurrency_limit': concurrency_count}
|
||||
demo.queue(**queue_kwargs).launch(
|
||||
server_name=args.server_name, server_port=args.server_port, share=args.share)
|
||||
|
||||
|
||||
def app_main(args: Optional[Union[List[str], AppArguments]] = None):
|
||||
return SwiftApp(args).main()
|
||||
@@ -0,0 +1,137 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import gradio as gr
|
||||
from functools import partial
|
||||
from typing import Literal, Optional
|
||||
|
||||
from swift.infer_engine import InferClient, InferRequest, RequestConfig
|
||||
from swift.template import History
|
||||
from swift.utils import get_file_mm_type
|
||||
from .locale import locale_mapping
|
||||
|
||||
|
||||
def clear_session():
|
||||
return '', [], []
|
||||
|
||||
|
||||
def modify_system_session(system: str):
|
||||
system = system or ''
|
||||
return system, '', [], []
|
||||
|
||||
|
||||
def _history_to_messages(history: History, system: Optional[str]):
|
||||
messages = []
|
||||
if system is not None:
|
||||
messages.append({'role': 'system', 'content': system})
|
||||
content = []
|
||||
for h in history:
|
||||
assert isinstance(h, (list, tuple))
|
||||
if isinstance(h[0], tuple):
|
||||
assert h[1] is None
|
||||
file_path = h[0][0]
|
||||
try:
|
||||
mm_type = get_file_mm_type(file_path)
|
||||
content.append({'type': mm_type, mm_type: file_path})
|
||||
except ValueError:
|
||||
with open(file_path, 'r', encoding='utf-8') as f:
|
||||
content.append({'type': 'text', 'text': f.read()})
|
||||
else:
|
||||
content.append({'type': 'text', 'text': h[0]})
|
||||
messages.append({'role': 'user', 'content': content})
|
||||
if h[1] is not None:
|
||||
messages.append({'role': 'assistant', 'content': h[1]})
|
||||
content = []
|
||||
return messages
|
||||
|
||||
|
||||
def _parse_text(text: str) -> str:
|
||||
mapping = {'<': '<', '>': '>', '*': '*'}
|
||||
for k, v in mapping.items():
|
||||
text = text.replace(k, v)
|
||||
return text
|
||||
|
||||
|
||||
async def model_chat(history: History, real_history: History, system: Optional[str], *, client, model: str,
|
||||
request_config: Optional[RequestConfig]):
|
||||
if history:
|
||||
|
||||
messages = _history_to_messages(real_history, system)
|
||||
resp_or_gen = await client.infer_async(
|
||||
InferRequest(messages=messages), request_config=request_config, model=model)
|
||||
if request_config and request_config.stream:
|
||||
response = ''
|
||||
async for resp in resp_or_gen:
|
||||
if resp is None:
|
||||
continue
|
||||
response += resp.choices[0].delta.content
|
||||
history[-1][1] = _parse_text(response)
|
||||
real_history[-1][-1] = response
|
||||
yield history, real_history
|
||||
|
||||
else:
|
||||
response = resp_or_gen.choices[0].message.content
|
||||
history[-1][1] = _parse_text(response)
|
||||
real_history[-1][-1] = response
|
||||
yield history, real_history
|
||||
|
||||
else:
|
||||
yield [], []
|
||||
|
||||
|
||||
def add_text(history: History, real_history: History, query: str):
|
||||
history = history or []
|
||||
real_history = real_history or []
|
||||
history.append([_parse_text(query), None])
|
||||
real_history.append([query, None])
|
||||
return history, real_history, ''
|
||||
|
||||
|
||||
def add_file(history: History, real_history: History, file):
|
||||
history = history or []
|
||||
real_history = real_history or []
|
||||
history.append([(file.name, ), None])
|
||||
real_history.append([(file.name, ), None])
|
||||
return history, real_history
|
||||
|
||||
|
||||
def build_ui(base_url: str,
|
||||
model: Optional[str] = None,
|
||||
*,
|
||||
request_config: Optional[RequestConfig] = None,
|
||||
is_multimodal: bool = True,
|
||||
studio_title: Optional[str] = None,
|
||||
lang: Literal['en', 'zh'] = 'en',
|
||||
default_system: Optional[str] = None):
|
||||
client = InferClient(base_url=base_url)
|
||||
model = model or client.models[0]
|
||||
studio_title = studio_title or model
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown(f'<center><font size=8>{studio_title}</center>')
|
||||
with gr.Row():
|
||||
with gr.Column(scale=3):
|
||||
system_input = gr.Textbox(value=default_system, lines=1, label='System')
|
||||
with gr.Column(scale=1):
|
||||
modify_system = gr.Button(locale_mapping['modify_system'][lang], scale=2)
|
||||
chatbot = gr.Chatbot(label='Chatbot')
|
||||
textbox = gr.Textbox(lines=1, label='Input')
|
||||
|
||||
with gr.Row():
|
||||
upload = gr.UploadButton(locale_mapping['upload'][lang], visible=is_multimodal)
|
||||
submit = gr.Button(locale_mapping['submit'][lang])
|
||||
regenerate = gr.Button(locale_mapping['regenerate'][lang])
|
||||
clear_history = gr.Button(locale_mapping['clear_history'][lang])
|
||||
|
||||
system_state = gr.State(value=default_system)
|
||||
history_state = gr.State(value=[])
|
||||
model_chat_ = partial(model_chat, client=client, model=model, request_config=request_config)
|
||||
|
||||
upload.upload(add_file, [chatbot, history_state, upload], [chatbot, history_state])
|
||||
textbox.submit(add_text, [chatbot, history_state, textbox],
|
||||
[chatbot, history_state, textbox]).then(model_chat_, [chatbot, history_state, system_state],
|
||||
[chatbot, history_state])
|
||||
submit.click(add_text, [chatbot, history_state, textbox],
|
||||
[chatbot, history_state, textbox]).then(model_chat_, [chatbot, history_state, system_state],
|
||||
[chatbot, history_state])
|
||||
regenerate.click(model_chat_, [chatbot, history_state, system_state], [chatbot, history_state])
|
||||
clear_history.click(clear_session, [], [textbox, chatbot, history_state])
|
||||
modify_system.click(modify_system_session, [system_input], [system_state, textbox, chatbot, history_state])
|
||||
return demo
|
||||
@@ -0,0 +1,23 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
locale_mapping = {
|
||||
'modify_system': {
|
||||
'en': '🛠️ Set system and clear history',
|
||||
'zh': '🛠️ 设置system并清空历史'
|
||||
},
|
||||
'clear_history': {
|
||||
'en': '🧹 Clear history',
|
||||
'zh': '🧹 清空历史'
|
||||
},
|
||||
'submit': {
|
||||
'en': '🚀 Send',
|
||||
'zh': '🚀 发送'
|
||||
},
|
||||
'regenerate': {
|
||||
'en': '🤔️ Regenerate',
|
||||
'zh': '🤔️ 重试'
|
||||
},
|
||||
'upload': {
|
||||
'en': '📁 Upload',
|
||||
'zh': '📁 上传'
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import datetime as dt
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import swift
|
||||
from swift.arguments import AppArguments, BaseArguments, WebUIArguments
|
||||
from swift.utils import ProcessorMixin, get_logger, parse_args, seed_everything
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftPipeline(ABC, ProcessorMixin):
|
||||
args_class = BaseArguments
|
||||
|
||||
def __init__(self, args: Optional[Union[List[str], args_class]] = None):
|
||||
self.args = self._parse_args(args)
|
||||
args = self.args
|
||||
logger.info(f'args: {args}')
|
||||
self._set_seed()
|
||||
self._compat_dsw_gradio(args)
|
||||
|
||||
def _set_seed(self):
|
||||
args = self.args
|
||||
if hasattr(args, 'seed'):
|
||||
seed = args.seed + max(getattr(args, 'rank', -1), 0)
|
||||
seed_everything(seed)
|
||||
logger.info(f'Global seed set to {seed}')
|
||||
|
||||
def _parse_args(self, args: Optional[Union[List[str], args_class]] = None) -> args_class:
|
||||
if isinstance(args, self.args_class):
|
||||
return args
|
||||
assert self.args_class is not None
|
||||
args, remaining_argv = parse_args(self.args_class, args)
|
||||
if len(remaining_argv) > 0:
|
||||
if getattr(args, 'ignore_args_error', False):
|
||||
logger.warning(f'remaining_argv: {remaining_argv}')
|
||||
else:
|
||||
raise ValueError(f'remaining_argv: {remaining_argv}')
|
||||
return args
|
||||
|
||||
@staticmethod
|
||||
def _compat_dsw_gradio(args) -> None:
|
||||
if (isinstance(args, (WebUIArguments, AppArguments)) and 'JUPYTER_NAME' in os.environ
|
||||
and 'dsw-' in os.environ['JUPYTER_NAME'] and 'GRADIO_ROOT_PATH' not in os.environ):
|
||||
os.environ['GRADIO_ROOT_PATH'] = f"/{os.environ['JUPYTER_NAME']}/proxy/{args.server_port}"
|
||||
|
||||
def main(self):
|
||||
logger.info(f'Start time of running main: {dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")}')
|
||||
logger.info(f'swift.__version__: {swift.__version__}')
|
||||
result = self.run()
|
||||
logger.info(f'End time of running main: {dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")}')
|
||||
return result
|
||||
|
||||
@abstractmethod
|
||||
def run(self):
|
||||
pass
|
||||
@@ -0,0 +1,2 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
from .eval import SwiftEval, eval_main
|
||||
@@ -0,0 +1,160 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import os
|
||||
from contextlib import nullcontext
|
||||
from evalscope.constants import EvalBackend, EvalType
|
||||
from evalscope.run import TaskConfig, run_task
|
||||
from evalscope.summarizer import Summarizer
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import EvalArguments
|
||||
from swift.dataset import MediaResource
|
||||
from swift.utils import append_to_jsonl, get_logger
|
||||
from ..base import SwiftPipeline
|
||||
from ..infer import run_deploy
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftEval(SwiftPipeline):
|
||||
args_class = EvalArguments
|
||||
args: args_class
|
||||
|
||||
def run(self):
|
||||
args = self.args
|
||||
eval_report = {}
|
||||
deploy_context = nullcontext() if args.eval_url else run_deploy(args, return_url=True)
|
||||
with deploy_context as base_url:
|
||||
base_url = args.eval_url or base_url
|
||||
|
||||
task_cfg = self.get_task_cfg(args.eval_dataset, args.eval_backend, base_url)
|
||||
result = self.get_task_result(task_cfg)
|
||||
eval_report[args.eval_backend] = result
|
||||
|
||||
eval_report.update({
|
||||
'time': args.time,
|
||||
'model': args.model,
|
||||
'adapters': args.adapters,
|
||||
'result_path': args.result_path,
|
||||
'eval_output_dir': args.eval_output_dir,
|
||||
'eval_limit': args.eval_limit
|
||||
})
|
||||
|
||||
if args.result_jsonl:
|
||||
append_to_jsonl(args.result_jsonl, eval_report)
|
||||
logger.info(f'The eval result have been saved to result_jsonl: `{args.result_jsonl}`.')
|
||||
return eval_report
|
||||
|
||||
def get_task_result(self, task_cfg: TaskConfig):
|
||||
run_task(task_cfg=task_cfg)
|
||||
reports = Summarizer.get_report_from_cfg(task_cfg=task_cfg)
|
||||
result = {}
|
||||
if task_cfg.eval_backend == EvalBackend.OPEN_COMPASS:
|
||||
for report in reports:
|
||||
if report[self.args.model_suffix] != '-':
|
||||
result[report['dataset']] = {report['metric']: report[self.args.model_suffix]}
|
||||
elif task_cfg.eval_backend == EvalBackend.VLM_EVAL_KIT:
|
||||
for report in reports:
|
||||
splited_key = next(iter(report)).rsplit('_', 2)
|
||||
if len(splited_key) == 3:
|
||||
_, dataset, metric = splited_key
|
||||
else:
|
||||
dataset, metric = '-', '-'
|
||||
result[dataset] = {metric: list(report.values())[0]}
|
||||
else:
|
||||
result = reports
|
||||
return result
|
||||
|
||||
def get_task_cfg(self, dataset: List[str], eval_backend: str, url: str):
|
||||
assert eval_backend in {EvalBackend.NATIVE, EvalBackend.OPEN_COMPASS, EvalBackend.VLM_EVAL_KIT}
|
||||
if eval_backend == EvalBackend.OPEN_COMPASS:
|
||||
if self.args.local_dataset:
|
||||
if os.path.exists('data'):
|
||||
if not os.path.exists(os.path.join('data', 'CMB')):
|
||||
raise RuntimeError('Opencompass need a `data` folder in your work dir('
|
||||
'which will be created automatically by swift eval), '
|
||||
'but a local path named `data` already exists, '
|
||||
'please consider moving the dir to another location.')
|
||||
else:
|
||||
local_dir = MediaResource.download(
|
||||
'https://modelscope.cn/datasets/'
|
||||
'opencompass/OpenCompassDataComplete/'
|
||||
'resolve/master/OpenCompassData-complete-20240207.zip', 'OpenCompassData')
|
||||
os.symlink(os.path.join(local_dir, 'data'), 'data')
|
||||
|
||||
task_cfg = self.get_opencompass_task_cfg(dataset, url)
|
||||
elif eval_backend == EvalBackend.VLM_EVAL_KIT:
|
||||
task_cfg = self.get_vlmeval_task_cfg(dataset, url)
|
||||
else:
|
||||
task_cfg = self.get_native_task_cfg(dataset, url)
|
||||
return task_cfg
|
||||
|
||||
def get_native_task_cfg(self, dataset: List[str], url: str):
|
||||
args = self.args
|
||||
work_dir = os.path.join(args.eval_output_dir, 'native')
|
||||
return TaskConfig(
|
||||
model=args.model_suffix,
|
||||
eval_type=EvalType.SERVICE,
|
||||
api_url=url,
|
||||
api_key=args.api_key or 'EMPTY',
|
||||
datasets=dataset,
|
||||
work_dir=work_dir,
|
||||
limit=args.eval_limit,
|
||||
eval_batch_size=args.eval_num_proc,
|
||||
dataset_args=args.eval_dataset_args,
|
||||
generation_config=args.eval_generation_config,
|
||||
**args.extra_eval_args)
|
||||
|
||||
def get_opencompass_task_cfg(self, dataset: List[str], url: str):
|
||||
# Must use chat/completion endpoint
|
||||
url = f"{url.rstrip('/')}/chat/completions"
|
||||
|
||||
args = self.args
|
||||
work_dir = os.path.join(args.eval_output_dir, 'opencompass')
|
||||
return TaskConfig(
|
||||
eval_backend=EvalBackend.OPEN_COMPASS,
|
||||
eval_config={
|
||||
'datasets':
|
||||
dataset,
|
||||
'batch_size':
|
||||
args.eval_num_proc,
|
||||
'work_dir':
|
||||
work_dir,
|
||||
'models': [{
|
||||
'path': args.model_suffix,
|
||||
'openai_api_base': url,
|
||||
'key': args.api_key or 'EMPTY',
|
||||
'is_chat': args.use_chat_template
|
||||
}],
|
||||
'limit':
|
||||
args.eval_limit
|
||||
},
|
||||
work_dir=work_dir)
|
||||
|
||||
def get_vlmeval_task_cfg(self, dataset: List[str], url: str):
|
||||
# Must use chat/completion endpoint
|
||||
url = f"{url.rstrip('/')}/chat/completions"
|
||||
|
||||
args = self.args
|
||||
work_dir = os.path.join(args.eval_output_dir, 'vlmeval')
|
||||
return TaskConfig(
|
||||
eval_backend=EvalBackend.VLM_EVAL_KIT,
|
||||
eval_config={
|
||||
'data':
|
||||
dataset,
|
||||
'model': [{
|
||||
'type': args.model_suffix,
|
||||
'name': 'CustomAPIModel',
|
||||
'api_base': url,
|
||||
'key': args.api_key or 'EMPTY',
|
||||
**args.eval_generation_config
|
||||
}],
|
||||
'nproc':
|
||||
args.eval_num_proc,
|
||||
'limit':
|
||||
args.eval_limit
|
||||
},
|
||||
work_dir=work_dir)
|
||||
|
||||
|
||||
def eval_main(args: Optional[Union[List[str], EvalArguments]] = None):
|
||||
return SwiftEval(args).main()
|
||||
@@ -0,0 +1,267 @@
|
||||
"""
|
||||
EvalScope integration utilities for ms-swift models.
|
||||
|
||||
This module provides a custom ModelAPI implementation that enables batch inference
|
||||
for evaluation tasks using ms-swift's TransformersEngine. It implements an asynchronous
|
||||
batch processing system to improve throughput when evaluating models.
|
||||
"""
|
||||
|
||||
from concurrent.futures import Future
|
||||
from dataclasses import dataclass
|
||||
from evalscope.api.messages import ChatMessage as EvalChatMessage
|
||||
from evalscope.api.model import GenerateConfig, ModelAPI, ModelOutput, ModelUsage
|
||||
from evalscope.api.registry import register_model_api
|
||||
from evalscope.api.tool import ToolChoice, ToolInfo
|
||||
from evalscope.models.utils.openai import chat_choices_from_openai
|
||||
from queue import Empty, Queue
|
||||
from threading import Thread
|
||||
from typing import Any, List, Optional, Tuple
|
||||
|
||||
from swift.infer_engine import InferRequest, RequestConfig, TransformersEngine
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchInferInput:
|
||||
"""
|
||||
Container for batch inference input data.
|
||||
|
||||
Holds all necessary information for a single inference request
|
||||
that will be processed as part of a batch.
|
||||
"""
|
||||
ms_input: InferRequest # ms-swift format request
|
||||
ms_config: RequestConfig # ms-swift format configuration
|
||||
batch_size: int # desired batch size for this request
|
||||
engine: TransformersEngine # inference engine to use
|
||||
|
||||
|
||||
@dataclass
|
||||
class _QueueItem:
|
||||
"""
|
||||
Internal queue item for batch processing.
|
||||
|
||||
Pairs a batch input with its corresponding future for result delivery.
|
||||
"""
|
||||
input: BatchInferInput
|
||||
future: Future[ModelOutput] # will be resolved with the inference result
|
||||
|
||||
|
||||
# Global variables for batch processing
|
||||
# These maintain the shared batch processing infrastructure across all model instances
|
||||
batch_thread: Optional[Thread] = None # background thread for processing batches
|
||||
batch_queue: Queue[_QueueItem] = Queue() # queue of pending inference requests
|
||||
|
||||
|
||||
@register_model_api('swift_custom')
|
||||
class EvalModel(ModelAPI):
|
||||
"""
|
||||
Custom ModelAPI implementation for ms-swift models with batch inference support.
|
||||
|
||||
This class integrates ms-swift's TransformersEngine with EvalScope's evaluation framework,
|
||||
providing efficient batch processing for improved evaluation throughput.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name: str,
|
||||
base_url: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
config: GenerateConfig = GenerateConfig(),
|
||||
**model_args: Any,
|
||||
):
|
||||
"""
|
||||
Initialize the EvalModel with ms-swift backend.
|
||||
|
||||
Args:
|
||||
model_name: Name of the model for identification
|
||||
base_url: Not used in this implementation (for API compatibility)
|
||||
api_key: Not used in this implementation (for API compatibility)
|
||||
config: Generation configuration with batch settings
|
||||
**model_args: Additional arguments including 'model' and 'template'
|
||||
"""
|
||||
super().__init__(
|
||||
model_name=model_name,
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
config=config,
|
||||
)
|
||||
|
||||
# Extract model-specific arguments from kwargs
|
||||
# This pattern allows us to collect known arguments while preserving unknown ones
|
||||
def collect_model_arg(name: str) -> Optional[Any]:
|
||||
value = model_args.get(name, None)
|
||||
if value is not None:
|
||||
model_args.pop(name)
|
||||
return value
|
||||
|
||||
# Extract required model parameters
|
||||
self.model = collect_model_arg('model') # model path or identifier
|
||||
self.template = collect_model_arg('template') # conversation template
|
||||
self.max_batch_size = collect_model_arg('max_batch_size') # maximum batch size
|
||||
|
||||
# Initialize the inference engine with batch support
|
||||
self.engine = TransformersEngine(self.model, template=self.template, max_batch_size=self.max_batch_size)
|
||||
|
||||
def generate(
|
||||
self,
|
||||
input: List[EvalChatMessage],
|
||||
tools: List[ToolInfo],
|
||||
tool_choice: ToolChoice,
|
||||
config: GenerateConfig,
|
||||
) -> ModelOutput:
|
||||
"""
|
||||
Generate model response using batch inference.
|
||||
|
||||
This method queues the request for batch processing and waits for the result.
|
||||
The actual inference is performed asynchronously in a background thread.
|
||||
|
||||
Args:
|
||||
input: List of chat messages forming the conversation
|
||||
tools: Available tools for function calling (if supported)
|
||||
tool_choice: Tool selection strategy
|
||||
config: Generation configuration
|
||||
|
||||
Returns:
|
||||
ModelOutput containing the generated response
|
||||
"""
|
||||
# Ensure the background batch processing thread is running
|
||||
global batch_thread
|
||||
if batch_thread is None:
|
||||
batch_thread = Thread(target=_process_batches, daemon=True)
|
||||
batch_thread.start()
|
||||
|
||||
# Convert EvalScope format to ms-swift format
|
||||
ms_input = convert_request(input, tools)
|
||||
ms_config = convert_config(config)
|
||||
|
||||
# Package the request for batch processing
|
||||
batch_input = BatchInferInput(
|
||||
ms_input=ms_input, ms_config=ms_config, batch_size=config.batch_size, engine=self.engine)
|
||||
|
||||
# Create a future to receive the result asynchronously
|
||||
future = Future[ModelOutput]()
|
||||
|
||||
# Queue the request for batch processing
|
||||
batch_queue.put(_QueueItem(input=batch_input, future=future))
|
||||
|
||||
# Block until the result is available
|
||||
return future.result()
|
||||
|
||||
|
||||
def _process_batches() -> None:
|
||||
"""
|
||||
Background thread function that processes batched inference requests.
|
||||
|
||||
This function runs continuously, collecting requests from the queue and
|
||||
processing them in batches for improved efficiency. It uses a timeout-based
|
||||
approach to balance between batch size and latency.
|
||||
"""
|
||||
while True:
|
||||
# Collect requests from the queue until timeout or batch size limit
|
||||
inputs: List[Tuple[BatchInferInput, Future[ModelOutput]]] = []
|
||||
|
||||
while True:
|
||||
try:
|
||||
# Wait for new requests with a 2-second timeout
|
||||
item = batch_queue.get(timeout=2)
|
||||
inputs.append((item.input, item.future))
|
||||
|
||||
# Check if we've reached the desired batch size
|
||||
if len(inputs) == item.input.batch_size:
|
||||
break # Process this batch now
|
||||
|
||||
except Empty:
|
||||
# No more requests in queue, process what we have
|
||||
break
|
||||
|
||||
# Skip processing if no requests were collected
|
||||
if len(inputs) == 0:
|
||||
continue
|
||||
|
||||
try:
|
||||
# Prepare batch inputs for ms-swift inference
|
||||
ms_inputs = [item[0].ms_input for item in inputs]
|
||||
ms_config = inputs[0][0].ms_config # use first config for the batch
|
||||
engine = inputs[0][0].engine # use first engine for the batch
|
||||
|
||||
# Perform batch inference using ms-swift engine
|
||||
completions = engine.infer(ms_inputs, ms_config, use_tqdm=False)
|
||||
|
||||
# Process results and deliver them to waiting futures
|
||||
for i, (batch_input, future) in enumerate(inputs):
|
||||
completion = completions[i]
|
||||
|
||||
# Convert ms-swift response to EvalScope format
|
||||
choices = chat_choices_from_openai(completion, tools=[])
|
||||
result = ModelOutput(
|
||||
model=completion.model,
|
||||
choices=choices,
|
||||
usage=(ModelUsage(
|
||||
input_tokens=completion.usage.prompt_tokens,
|
||||
output_tokens=completion.usage.completion_tokens,
|
||||
total_tokens=completion.usage.total_tokens,
|
||||
) if completion.usage else None),
|
||||
)
|
||||
|
||||
# Deliver the result to the waiting caller
|
||||
future.set_result(result)
|
||||
|
||||
except Exception as ex:
|
||||
# If batch processing fails, propagate the error to all waiting futures
|
||||
for _, future in inputs:
|
||||
future.set_exception(ex)
|
||||
|
||||
|
||||
def convert_config(config: GenerateConfig) -> RequestConfig:
|
||||
"""
|
||||
Convert EvalScope GenerateConfig to ms-swift RequestConfig.
|
||||
|
||||
Maps configuration parameters between the two frameworks, ensuring
|
||||
compatibility while maintaining the same generation behavior.
|
||||
|
||||
Args:
|
||||
config: EvalScope generation configuration
|
||||
|
||||
Returns:
|
||||
RequestConfig: ms-swift compatible configuration
|
||||
"""
|
||||
return RequestConfig(
|
||||
max_tokens=config.max_tokens,
|
||||
temperature=config.temperature,
|
||||
top_k=config.top_k,
|
||||
top_p=config.top_p,
|
||||
presence_penalty=config.presence_penalty,
|
||||
frequency_penalty=config.frequency_penalty,
|
||||
seed=config.seed,
|
||||
stream=False, # batch processing doesn't support streaming
|
||||
logprobs=config.logprobs,
|
||||
top_logprobs=config.top_logprobs)
|
||||
|
||||
|
||||
def convert_request(messages: List[EvalChatMessage], tools: List[ToolInfo]) -> InferRequest:
|
||||
"""
|
||||
Convert EvalScope request format to ms-swift InferRequest format.
|
||||
|
||||
Transforms the message and tool format from EvalScope's representation
|
||||
to the format expected by ms-swift's inference engine.
|
||||
|
||||
Args:
|
||||
messages: List of chat messages in EvalScope format
|
||||
tools: List of available tools in EvalScope format
|
||||
|
||||
Returns:
|
||||
InferRequest: ms-swift compatible request object
|
||||
"""
|
||||
# Convert tools to ms-swift format
|
||||
tools_list = []
|
||||
if len(tools) > 0:
|
||||
tools_list = [tool.model_dump(exclude_none=True) for tool in tools]
|
||||
|
||||
# Convert messages to ms-swift format
|
||||
ms_messages = []
|
||||
for message in messages:
|
||||
ms_messages.append(message.model_dump(exclude_none=True))
|
||||
|
||||
return InferRequest(
|
||||
messages=ms_messages,
|
||||
tools=tools_list,
|
||||
)
|
||||
@@ -0,0 +1,6 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
from .cached_dataset import export_cached_dataset
|
||||
from .export import SwiftExport, export_main
|
||||
from .merge_lora import merge_lora
|
||||
from .ollama import export_to_ollama
|
||||
from .quant import quantize_model
|
||||
@@ -0,0 +1,47 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import os
|
||||
import torch
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import ExportArguments
|
||||
from swift.utils import get_logger
|
||||
from ..train import SwiftSft
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class ExportCachedDataset(SwiftSft):
|
||||
args_class = ExportArguments
|
||||
args: args_class
|
||||
|
||||
def __init__(self, args: Optional[Union[List[str], ExportArguments]] = None) -> None:
|
||||
super(SwiftSft, self).__init__(args)
|
||||
args = self.args
|
||||
self.train_msg = {} # dummy
|
||||
template_cls = args.template_meta.template_cls
|
||||
if template_cls and template_cls.use_model:
|
||||
kwargs = {'return_dummy_model': True}
|
||||
else:
|
||||
kwargs = {'load_model': False}
|
||||
with torch.device('meta'):
|
||||
self._prepare_model_tokenizer(**kwargs)
|
||||
self._prepare_template()
|
||||
self.template.set_mode(args.template_mode)
|
||||
|
||||
def _post_process_datasets(self, datasets: List) -> List:
|
||||
return datasets
|
||||
|
||||
def main(self):
|
||||
train_dataset, val_dataset = self._prepare_dataset()
|
||||
train_data_dir = os.path.join(self.args.output_dir, 'train')
|
||||
val_data_dir = os.path.join(self.args.output_dir, 'val')
|
||||
train_dataset.save_to_disk(train_data_dir)
|
||||
if val_dataset is not None:
|
||||
val_dataset.save_to_disk(val_data_dir)
|
||||
logger.info(f'cached_dataset: `{train_data_dir}`')
|
||||
if val_dataset is not None:
|
||||
logger.info(f'cached_val_dataset: `{val_data_dir}`')
|
||||
|
||||
|
||||
def export_cached_dataset(args: Optional[Union[List[str], ExportArguments]] = None):
|
||||
return ExportCachedDataset(args).main()
|
||||
@@ -0,0 +1,54 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import ExportArguments
|
||||
from swift.pipelines import SwiftPipeline
|
||||
from swift.tuners import swift_to_peft_format
|
||||
from swift.utils import get_logger
|
||||
from .cached_dataset import export_cached_dataset
|
||||
from .merge_lora import merge_lora
|
||||
from .ollama import export_to_ollama
|
||||
from .quant import quantize_model
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftExport(SwiftPipeline):
|
||||
args_class = ExportArguments
|
||||
args: args_class
|
||||
|
||||
def run(self):
|
||||
args = self.args
|
||||
if args.to_peft_format:
|
||||
args.adapters[0] = swift_to_peft_format(args.adapters[0], args.output_dir)
|
||||
if args.merge_lora:
|
||||
output_dir = args.output_dir
|
||||
if args.to_peft_format or args.quant_method or args.to_ollama or args.push_to_hub:
|
||||
args.output_dir = None
|
||||
merge_lora(args)
|
||||
args.output_dir = output_dir # recover
|
||||
if args.quant_method:
|
||||
quantize_model(args)
|
||||
elif args.to_ollama:
|
||||
export_to_ollama(args)
|
||||
elif args.to_cached_dataset:
|
||||
export_cached_dataset(args)
|
||||
elif args.to_hf or args.mcore_adapter and args.to_mcore:
|
||||
from swift.megatron import convert_mcore2hf
|
||||
convert_mcore2hf(args)
|
||||
elif args.to_mcore:
|
||||
from swift.megatron import convert_hf2mcore
|
||||
convert_hf2mcore(args)
|
||||
elif args.push_to_hub:
|
||||
model_dir = args.adapters and args.adapters[0] or args.model_dir
|
||||
assert model_dir, f'model_dir: {model_dir}'
|
||||
args.hub.push_to_hub(
|
||||
args.hub_model_id,
|
||||
model_dir,
|
||||
token=args.hub_token,
|
||||
private=args.hub_private_repo,
|
||||
commit_message=args.commit_message)
|
||||
|
||||
|
||||
def export_main(args: Optional[Union[List[str], ExportArguments]] = None):
|
||||
return SwiftExport(args).main()
|
||||
@@ -0,0 +1,61 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import os
|
||||
|
||||
from swift.arguments import ExportArguments
|
||||
from swift.model import save_checkpoint
|
||||
from swift.tuners import Swift
|
||||
from swift.utils import HfConfigFactory, get_logger
|
||||
from ..utils import prepare_model_template
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def check_tie_word_embeddings(model):
|
||||
config = model.config
|
||||
try:
|
||||
from peft.utils import ModulesToSaveWrapper
|
||||
if not HfConfigFactory.get_config_attr(config, 'tie_word_embeddings'):
|
||||
return
|
||||
for module in [model.get_input_embeddings(), model.get_output_embeddings()]:
|
||||
if not isinstance(module, ModulesToSaveWrapper):
|
||||
return
|
||||
HfConfigFactory.set_config_attr(config, 'tie_word_embeddings', False)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def merge_lora(args: ExportArguments, device_map=None, replace_if_exists=False) -> None:
|
||||
if replace_if_exists:
|
||||
logger.info(f'replace_if_exists: {replace_if_exists}')
|
||||
output_dir = getattr(args, 'output_dir', None) or f'{args.adapters[0]}-merged'
|
||||
if os.path.exists(output_dir) and not replace_if_exists:
|
||||
logger.info(f'The weight directory for the merged LoRA already exists in {output_dir}, '
|
||||
'skipping the saving process.')
|
||||
else:
|
||||
# If the model is quantized, perform the merge on the original (unquantized) model.
|
||||
# https://github.com/huggingface/peft/issues/2321
|
||||
args.quant_method = None
|
||||
origin_device_map = args.device_map
|
||||
args.device_map = device_map or args.device_map
|
||||
logger.info(f'merge_device_map: {device_map}')
|
||||
model, template = prepare_model_template(args)
|
||||
logger.info('Merge LoRA...')
|
||||
check_tie_word_embeddings(model)
|
||||
Swift.merge_and_unload(model)
|
||||
model = model.model
|
||||
logger.info('Saving merged weights...')
|
||||
|
||||
save_checkpoint(
|
||||
model,
|
||||
template.processor,
|
||||
output_dir,
|
||||
safe_serialization=args.safe_serialization,
|
||||
model_dirs=args.adapters,
|
||||
max_shard_size=args.max_shard_size,
|
||||
additional_saved_files=model.model_meta.additional_saved_files)
|
||||
logger.info(f'Successfully merged LoRA and saved in `{output_dir}`.')
|
||||
args.device_map = origin_device_map
|
||||
|
||||
args.model = output_dir
|
||||
args.model_dir = output_dir
|
||||
args.adapters = []
|
||||
@@ -0,0 +1,73 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import os
|
||||
from typing import List
|
||||
|
||||
from swift.arguments import ExportArguments
|
||||
from swift.infer_engine import RequestConfig, TransformersEngine
|
||||
from swift.template import Template
|
||||
from swift.utils import get_logger
|
||||
from ..utils import prepare_model_template
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def replace_and_concat(template: 'Template', template_list: List, placeholder: str, keyword: str):
|
||||
final_str = ''
|
||||
for t in template_list:
|
||||
if isinstance(t, str):
|
||||
final_str += t.replace(placeholder, keyword)
|
||||
elif isinstance(t, (tuple, list)):
|
||||
if isinstance(t[0], int):
|
||||
final_str += template.tokenizer.decode(t)
|
||||
else:
|
||||
for attr in t:
|
||||
if attr == 'bos_token_id':
|
||||
final_str += template.tokenizer.bos_token
|
||||
elif attr == 'eos_token_id':
|
||||
final_str += template.tokenizer.eos_token
|
||||
else:
|
||||
raise ValueError(f'Unknown token: {attr}')
|
||||
return final_str
|
||||
|
||||
|
||||
def export_to_ollama(args: ExportArguments):
|
||||
args.device_map = 'meta' # Accelerate load speed.
|
||||
logger.info('Exporting to ollama:')
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
model, template = prepare_model_template(args)
|
||||
engine = TransformersEngine(model, template=template)
|
||||
logger.info(f'Using model_dir: {engine.model_dir}')
|
||||
template_meta = template.template_meta
|
||||
with open(os.path.join(args.output_dir, 'Modelfile'), 'w', encoding='utf-8') as f:
|
||||
f.write(f'FROM {engine.model_dir}\n')
|
||||
f.write(f'TEMPLATE """{{{{ if .System }}}}'
|
||||
f'{replace_and_concat(template, template_meta.system_prefix, "{{SYSTEM}}", "{{ .System }}")}'
|
||||
f'{{{{ else }}}}{replace_and_concat(template, template_meta.prefix, "", "")}'
|
||||
f'{{{{ end }}}}')
|
||||
f.write(f'{{{{ if .Prompt }}}}'
|
||||
f'{replace_and_concat(template, template_meta.prompt, "{{QUERY}}", "{{ .Prompt }}")}'
|
||||
f'{{{{ end }}}}')
|
||||
f.write('{{ .Response }}')
|
||||
f.write(replace_and_concat(template, template_meta.suffix, '', '') + '"""\n')
|
||||
f.write(f'PARAMETER stop "{replace_and_concat(template, template_meta.suffix, "", "")}"\n')
|
||||
|
||||
request_config = RequestConfig(
|
||||
temperature=args.temperature,
|
||||
top_k=args.top_k,
|
||||
top_p=args.top_p,
|
||||
repetition_penalty=args.repetition_penalty)
|
||||
generation_config = engine._prepare_generation_config(request_config)
|
||||
engine._add_stop_words(generation_config, request_config)
|
||||
for stop_word in generation_config.stop_words:
|
||||
f.write(f'PARAMETER stop "{stop_word}"\n')
|
||||
f.write(f'PARAMETER temperature {generation_config.temperature}\n')
|
||||
f.write(f'PARAMETER top_k {generation_config.top_k}\n')
|
||||
f.write(f'PARAMETER top_p {generation_config.top_p}\n')
|
||||
f.write(f'PARAMETER repeat_penalty {generation_config.repetition_penalty}\n')
|
||||
|
||||
logger.info('Save Modelfile done, you can start ollama by:')
|
||||
logger.info('> ollama serve')
|
||||
logger.info('In another terminal:')
|
||||
logger.info('> ollama create my-custom-model '
|
||||
f'-f {os.path.join(args.output_dir, "Modelfile")}')
|
||||
logger.info('> ollama run my-custom-model')
|
||||
@@ -0,0 +1,292 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import transformers
|
||||
from collections import defaultdict
|
||||
from contextlib import contextmanager
|
||||
from packaging import version
|
||||
from tqdm import tqdm
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from swift.arguments import ExportArguments
|
||||
from swift.dataset import load_dataset
|
||||
from swift.model import save_checkpoint
|
||||
from swift.template import MaxLengthError
|
||||
from swift.utils import HfConfigFactory, ProcessorMixin, deep_getattr, get_logger, get_model_parameter_info, to_device
|
||||
from ..utils import prepare_model_template
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class QuantEngine(ProcessorMixin):
|
||||
|
||||
def __init__(self, args: ExportArguments):
|
||||
self.args = args
|
||||
kwargs = {}
|
||||
if args.quant_method == 'awq':
|
||||
from awq import AutoAWQForCausalLM
|
||||
kwargs['auto_model_cls'] = AutoAWQForCausalLM
|
||||
self.model, self.template = prepare_model_template(args, **kwargs)
|
||||
self.template.set_mode('train')
|
||||
self.model.config.use_cache = False
|
||||
HfConfigFactory.set_config_attr(self.model.config, 'use_cache', False)
|
||||
self.processor = self.template.processor
|
||||
args.save_args()
|
||||
|
||||
def quantize(self):
|
||||
args = self.args
|
||||
if args.quant_bits is None and args.quant_method != 'fp8':
|
||||
raise ValueError(f'Please set the quant_bits. args.quant_bits: {args.quant_bits}')
|
||||
if args.quant_method == 'awq':
|
||||
self.template.model = self.model.model
|
||||
self.awq_model_quantize()
|
||||
self.model.save_quantized(
|
||||
args.output_dir, safetensors=args.safe_serialization, shard_size=args.max_shard_size)
|
||||
elif args.quant_method in {'gptq', 'gptq_v2'}:
|
||||
self.template.model = self.model
|
||||
gptq_quantizer = self.gptq_model_quantize(v2=(args.quant_method == 'gptq_v2'))
|
||||
if args.quant_method == 'gptq_v2':
|
||||
if not getattr(self.model, '_dynamic_tied_weights_keys', None):
|
||||
self.model._dynamic_tied_weights_keys = []
|
||||
self.model._dynamic_tied_weights_keys += ['wf_unsqueeze_zero', 'wf_unsqueeze_neg_one']
|
||||
gptq_quantizer.save(
|
||||
self.model,
|
||||
args.output_dir,
|
||||
safe_serialization=args.safe_serialization,
|
||||
max_shard_size=args.max_shard_size)
|
||||
elif args.quant_method in {'bnb', 'fp8'}:
|
||||
self.model.save_pretrained(
|
||||
args.output_dir, safe_serialization=args.safe_serialization, max_shard_size=args.max_shard_size)
|
||||
else:
|
||||
raise ValueError(f'args.quant_method: {args.quant_method}')
|
||||
|
||||
logger.info(f'model: {self.model}')
|
||||
logger.info(f'model_parameter_info: {get_model_parameter_info(self.model)}')
|
||||
save_checkpoint(
|
||||
None,
|
||||
self.processor,
|
||||
args.output_dir,
|
||||
model_dirs=[args.model_dir],
|
||||
additional_saved_files=self.model.model_meta.additional_saved_files)
|
||||
logger.info(f'Successfully quantized the model and saved in `{args.output_dir}`.')
|
||||
|
||||
@torch.inference_mode()
|
||||
def _prepare_gptq_dataset(self, examples: List[Dict[str, torch.LongTensor]], batch_size: int = 1, *args, **kwargs):
|
||||
res = []
|
||||
for start in tqdm(range(0, len(examples), batch_size)):
|
||||
batched_inputs = examples[start:start + batch_size]
|
||||
inputs = to_device(self.template.data_collator(batched_inputs), self.model.device)
|
||||
if self.model.model_meta.is_multimodal:
|
||||
_, inputs = self.template.pre_forward_hook(self.model, None, inputs)
|
||||
res.append(to_device(inputs, 'cpu'))
|
||||
return res
|
||||
|
||||
@torch.inference_mode()
|
||||
def _get_quant_dataset(self, *args, **kwargs):
|
||||
args = self.args
|
||||
assert args.quant_method in {'awq', 'gptq', 'gptq_v2'}
|
||||
template = self.template
|
||||
n_samples = args.quant_n_samples
|
||||
block_size = args.max_length
|
||||
|
||||
# only use train_dataset
|
||||
dataset = load_dataset(
|
||||
args.dataset, split_dataset_ratio=0, shuffle=args.dataset_shuffle, **args.get_dataset_kwargs())[0]
|
||||
logger.info(f'quant_dataset: {dataset}')
|
||||
dataset = dataset.shuffle()
|
||||
|
||||
samples = []
|
||||
i = 0
|
||||
prog_bar = tqdm(total=n_samples, dynamic_ncols=True)
|
||||
is_multimodal = self.model.model_meta.is_multimodal
|
||||
for data in dataset:
|
||||
try:
|
||||
inputs = template.encode(data)
|
||||
except MaxLengthError:
|
||||
continue
|
||||
if is_multimodal and args.quant_method in {'gptq', 'gptq_v2'}:
|
||||
inputs.pop('labels', None)
|
||||
samples.append(inputs)
|
||||
else:
|
||||
input_ids = inputs['input_ids']
|
||||
samples += input_ids
|
||||
i += 1
|
||||
prog_bar.update()
|
||||
if i == n_samples:
|
||||
break
|
||||
prog_bar.close()
|
||||
if is_multimodal and args.quant_method in {'gptq', 'gptq_v2'}:
|
||||
return samples
|
||||
# now concatenate all samples and split according to block size
|
||||
n_split = max(len(samples) // block_size, 1)
|
||||
logger.info(f'Split into {n_split} blocks')
|
||||
res = []
|
||||
for i in range(n_split):
|
||||
input_ids = samples[i * block_size:(i + 1) * block_size]
|
||||
if args.quant_method in {'gptq', 'gptq_v2'}:
|
||||
res.append({'input_ids': input_ids})
|
||||
else:
|
||||
res.append(torch.tensor(input_ids)[None])
|
||||
return res
|
||||
|
||||
@staticmethod
|
||||
@contextmanager
|
||||
def _patch_awq_move_embed(awq_model):
|
||||
_origin_move_embed = awq_model.move_embed
|
||||
|
||||
def _move_embed(model, device: str):
|
||||
if hasattr(model, '_hf_hook') and device != 'cpu':
|
||||
return
|
||||
_origin_move_embed(model, device)
|
||||
|
||||
awq_model.move_embed = _move_embed
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
awq_model.move_embed = _origin_move_embed
|
||||
|
||||
def awq_model_quantize(self) -> None:
|
||||
from awq.quantize import quantizer
|
||||
|
||||
args = self.args
|
||||
logger.info(f'Quantization dataset: {args.dataset}')
|
||||
_origin_get_calib_dataset = quantizer.get_calib_dataset
|
||||
quantizer.get_calib_dataset = self._get_quant_dataset
|
||||
quant_config = {
|
||||
'zero_point': True,
|
||||
'q_group_size': args.group_size,
|
||||
'w_bit': args.quant_bits,
|
||||
'version': 'GEMM'
|
||||
}
|
||||
if self.model.model_info.is_moe_model:
|
||||
quant_config['modules_to_not_convert'] = self.args.get_modules_to_not_convert()
|
||||
logger.info(f'quant_config: {quant_config}')
|
||||
logger.info('Start quantizing the model...')
|
||||
with self._patch_awq_move_embed(self.model):
|
||||
self.model.quantize(
|
||||
self.tokenizer, quant_config=quant_config, n_parallel_calib_samples=args.quant_batch_size)
|
||||
quantizer.get_calib_dataset = _origin_get_calib_dataset # recover
|
||||
if self.model.quant_config.modules_to_not_convert:
|
||||
model_arch = args.model_meta.model_arch
|
||||
lm_head_key = getattr(model_arch, 'lm_head', None) or 'lm_head'
|
||||
if lm_head_key not in self.model.quant_config.modules_to_not_convert:
|
||||
self.model.quant_config.modules_to_not_convert.append(lm_head_key)
|
||||
|
||||
@contextmanager
|
||||
def _patch_gptq(self):
|
||||
from optimum.gptq import quantizer
|
||||
_get_dataset_origin = quantizer.get_dataset
|
||||
_prepare_dataset_origin = quantizer.prepare_dataset
|
||||
quantizer.get_dataset = self._get_quant_dataset
|
||||
quantizer.prepare_dataset = self._prepare_gptq_dataset
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
quantizer.get_dataset = _get_dataset_origin
|
||||
quantizer.prepare_dataset = _prepare_dataset_origin
|
||||
|
||||
@staticmethod
|
||||
def get_block_name_to_quantize(model: nn.Module) -> Optional[str]:
|
||||
model_arch = model.model_meta.model_arch
|
||||
prefix = ''
|
||||
if hasattr(model_arch, 'language_model'):
|
||||
language_model = [lm for lm in model_arch.language_model if not lm.endswith('lm_head')]
|
||||
assert len(language_model) == 1, f'model_arch.language_model: {language_model}'
|
||||
prefix = language_model[0]
|
||||
model = deep_getattr(model, prefix)
|
||||
|
||||
module_lists = []
|
||||
for n, m in model.named_modules():
|
||||
if (isinstance(m, (nn.ModuleList, nn.Sequential)) and len(m) >= 10
|
||||
and 'mlp' not in m[0].__class__.__name__.lower()): # fix moe
|
||||
module_lists.append((n, m))
|
||||
if module_lists:
|
||||
module_list = max(module_lists, key=lambda x: len(x[1]))
|
||||
return f'{prefix}.{module_list[0]}'.strip('.')
|
||||
|
||||
@staticmethod
|
||||
def _get_experts(block):
|
||||
for n, m in block.named_modules():
|
||||
if isinstance(m, (nn.ModuleList, nn.Sequential)):
|
||||
return n, m
|
||||
|
||||
@staticmethod
|
||||
def get_modules_in_block_to_quantize(model, block_name: str):
|
||||
if not model.model_info.is_moe_model:
|
||||
return
|
||||
from optimum.gptq.utils import get_layers
|
||||
|
||||
# Do not quantize the gate part.
|
||||
block = deep_getattr(model, block_name)[-1]
|
||||
prefix, experts = QuantEngine._get_experts(block)
|
||||
layers = get_layers(block)
|
||||
res = []
|
||||
experts = defaultdict(list)
|
||||
experts_idx = None
|
||||
for name, layer in layers.items():
|
||||
if model.model_info.model_type == 'qwen3_next' and name.startswith('self_attn.'):
|
||||
# ignore attn
|
||||
continue
|
||||
if name.startswith(prefix):
|
||||
suffix = name.rsplit('.', 1)[-1]
|
||||
experts[suffix].append(name)
|
||||
experts_idx = len(res)
|
||||
elif 'mlp.gate' not in name:
|
||||
res.append([name])
|
||||
res[experts_idx:experts_idx] = experts.values()
|
||||
return res
|
||||
|
||||
@contextmanager
|
||||
def _patch_gptq_block(self, model, block_name_to_quantize):
|
||||
if version.parse(transformers.__version__) < version.parse('4.54'):
|
||||
yield
|
||||
return
|
||||
# compat transformers>=4.54
|
||||
blocks = deep_getattr(model, block_name_to_quantize)
|
||||
hooks = []
|
||||
|
||||
def _to_tuple(module, input, output):
|
||||
if not isinstance(output, (list, tuple)):
|
||||
output = (output, )
|
||||
return output
|
||||
|
||||
for block in blocks:
|
||||
hooks.append(block.register_forward_hook(_to_tuple))
|
||||
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
for hook in hooks:
|
||||
hook.remove()
|
||||
|
||||
def gptq_model_quantize(self, v2: bool = False):
|
||||
from optimum.gptq import GPTQQuantizer
|
||||
args = self.args
|
||||
logger.info(f'Quantization dataset: {args.dataset}')
|
||||
block_name_to_quantize = self.get_block_name_to_quantize(self.model)
|
||||
modules_in_block_to_quantize = self.get_modules_in_block_to_quantize(self.model, block_name_to_quantize)
|
||||
logger.info(f'block_name_to_quantize: {block_name_to_quantize}')
|
||||
logger.info(f'modules_in_block_to_quantize: {modules_in_block_to_quantize}')
|
||||
with self._patch_gptq():
|
||||
gptq_quantizer = GPTQQuantizer(
|
||||
bits=args.quant_bits,
|
||||
group_size=args.group_size,
|
||||
dataset=','.join(args.dataset),
|
||||
batch_size=args.quant_batch_size,
|
||||
block_name_to_quantize=block_name_to_quantize,
|
||||
modules_in_block_to_quantize=modules_in_block_to_quantize,
|
||||
checkpoint_format='gptq_v2' if v2 else 'gptq')
|
||||
gptq_quantizer.serialization_keys.append('block_name_to_quantize')
|
||||
logger.info('Start quantizing the model...')
|
||||
logger.warning('The process of packing the model takes a long time and there is no progress bar. '
|
||||
'Please be patient and wait...')
|
||||
if not hasattr(self.model, 'hf_device_map'):
|
||||
self.model.hf_device_map = {'': torch.device('cuda:0')}
|
||||
with self._patch_gptq_block(self.model, block_name_to_quantize):
|
||||
gptq_quantizer.quantize_model(self.model, self.tokenizer)
|
||||
self.model.config.quantization_config.pop('dataset', None)
|
||||
return gptq_quantizer
|
||||
|
||||
|
||||
def quantize_model(args: ExportArguments):
|
||||
QuantEngine(args).quantize()
|
||||
@@ -0,0 +1,26 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from swift.utils.import_utils import _LazyModule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .deploy import SwiftDeploy, deploy_main, run_deploy
|
||||
from .infer import SwiftInfer, infer_main
|
||||
from .rollout import rollout_main
|
||||
else:
|
||||
_import_structure = {
|
||||
'rollout': ['rollout_main'],
|
||||
'infer': ['infer_main', 'SwiftInfer'],
|
||||
'deploy': ['deploy_main', 'SwiftDeploy', 'run_deploy'],
|
||||
'protocol': ['RequestConfig', 'Function'],
|
||||
}
|
||||
|
||||
import sys
|
||||
|
||||
sys.modules[__name__] = _LazyModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
@@ -0,0 +1,289 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
import multiprocessing
|
||||
import time
|
||||
import uvicorn
|
||||
from aiohttp import ClientConnectorError
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import asdict
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
from http import HTTPStatus
|
||||
from threading import Thread
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import DeployArguments, InferArguments
|
||||
from swift.infer_engine import AdapterRequest, InferClient, RequestConfig
|
||||
from swift.infer_engine.protocol import (ChatCompletionRequest, CompletionRequest, EmbeddingRequest, Model, ModelList,
|
||||
MultiModalRequestMixin, RolloutInferRequest)
|
||||
from swift.metrics import InferStats
|
||||
from swift.utils import JsonlWriter, get_logger
|
||||
from .infer import SwiftInfer
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftDeploy(SwiftInfer):
|
||||
args_class = DeployArguments
|
||||
args: args_class
|
||||
|
||||
@staticmethod
|
||||
def get_infer_engine(args: InferArguments, template=None, **kwargs):
|
||||
if isinstance(args, DeployArguments) and args.infer_backend == 'vllm':
|
||||
engine_kwargs = (kwargs.get('engine_kwargs') or {}).copy()
|
||||
if args.vllm_data_parallel_size > 1:
|
||||
if not args.vllm_use_async_engine:
|
||||
raise ValueError('vLLM data parallel requires `vllm_use_async_engine=True` in deploy mode.')
|
||||
engine_kwargs.setdefault('data_parallel_size', args.vllm_data_parallel_size)
|
||||
logger.info(f'Enable vLLM data parallel with size {args.vllm_data_parallel_size}.')
|
||||
if args.max_logprobs is not None:
|
||||
engine_kwargs['max_logprobs'] = args.max_logprobs
|
||||
kwargs['engine_kwargs'] = engine_kwargs
|
||||
return SwiftInfer.get_infer_engine(args, template, **kwargs)
|
||||
|
||||
def _register_app(self):
|
||||
self.app.get('/health')(self.health)
|
||||
self.app.get('/ping')(self.ping)
|
||||
self.app.post('/ping')(self.ping)
|
||||
self.app.get('/v1/models')(self.get_available_models)
|
||||
self.app.post('/v1/chat/completions')(self.create_chat_completion)
|
||||
self.app.post('/v1/completions')(self.create_completion)
|
||||
self.app.post('/v1/embeddings')(self.create_embedding)
|
||||
self.app.post('/infer/')(self.infer_handler)
|
||||
|
||||
def __init__(self, args: Optional[Union[List[str], DeployArguments]] = None) -> None:
|
||||
super().__init__(args)
|
||||
|
||||
self.infer_engine.strict = True
|
||||
self.infer_stats = InferStats()
|
||||
self.app = FastAPI(lifespan=self.lifespan)
|
||||
self._register_app()
|
||||
|
||||
async def _log_stats_hook(self):
|
||||
while True:
|
||||
await asyncio.sleep(self.args.log_interval)
|
||||
self._compute_infer_stats()
|
||||
self.infer_stats.reset()
|
||||
|
||||
def _compute_infer_stats(self):
|
||||
global_stats = self.infer_stats.compute()
|
||||
for k, v in global_stats.items():
|
||||
global_stats[k] = round(v, 8)
|
||||
logger.info(global_stats)
|
||||
|
||||
def lifespan(self, app: FastAPI):
|
||||
args = self.args
|
||||
if args.log_interval > 0:
|
||||
thread = Thread(target=lambda: asyncio.run(self._log_stats_hook()), daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if args.log_interval > 0:
|
||||
self._compute_infer_stats()
|
||||
|
||||
def _get_model_list(self):
|
||||
args = self.args
|
||||
model_list = [args.served_model_name or args.model_suffix]
|
||||
if args.adapter_mapping:
|
||||
model_list += [name for name in args.adapter_mapping.keys()]
|
||||
return model_list
|
||||
|
||||
async def health(self) -> Response:
|
||||
"""Health check endpoint."""
|
||||
if self.infer_engine is not None:
|
||||
return Response(status_code=200)
|
||||
else:
|
||||
return Response(status_code=503)
|
||||
|
||||
async def ping(self) -> Response:
|
||||
"""Ping check endpoint. Required for SageMaker compatibility."""
|
||||
return await self.health()
|
||||
|
||||
async def get_available_models(self):
|
||||
model_list = self._get_model_list()
|
||||
data = [Model(id=model_id, owned_by=self.args.owned_by) for model_id in model_list]
|
||||
return ModelList(data=data)
|
||||
|
||||
async def _check_model(self, request: ChatCompletionRequest) -> Optional[str]:
|
||||
available_models = await self.get_available_models()
|
||||
model_list = [model.id for model in available_models.data]
|
||||
if request.model not in model_list:
|
||||
return f'`{request.model}` is not in the model_list: `{model_list}`.'
|
||||
|
||||
def _check_api_key(self, raw_request: Request) -> Optional[str]:
|
||||
api_key = self.args.api_key
|
||||
if api_key is None:
|
||||
return
|
||||
authorization = dict(raw_request.headers).get('authorization')
|
||||
error_msg = 'API key error'
|
||||
if authorization is None or not authorization.startswith('Bearer '):
|
||||
return error_msg
|
||||
request_api_key = authorization[7:]
|
||||
if request_api_key != api_key:
|
||||
return error_msg
|
||||
|
||||
def _check_max_logprobs(self, request):
|
||||
args = self.args
|
||||
if isinstance(request.top_logprobs, int) and request.top_logprobs > args.max_logprobs:
|
||||
return (f'The value of top_logprobs({request.top_logprobs}) is greater than '
|
||||
f'the server\'s max_logprobs({args.max_logprobs}).')
|
||||
|
||||
@staticmethod
|
||||
def create_error_response(status_code: Union[int, str, HTTPStatus], message: str) -> JSONResponse:
|
||||
status_code = int(status_code)
|
||||
return JSONResponse({'message': message, 'object': 'error'}, status_code)
|
||||
|
||||
def _post_process(self, request_info, response, return_cmpl_response: bool = False):
|
||||
args = self.args
|
||||
|
||||
for i in range(len(response.choices)):
|
||||
if not hasattr(response.choices[i], 'message') or not isinstance(response.choices[i].message.content,
|
||||
(tuple, list)):
|
||||
continue
|
||||
for j, content in enumerate(response.choices[i].message.content):
|
||||
if isinstance(content, dict) and content['type'] == 'image':
|
||||
b64_image = MultiModalRequestMixin.to_base64(content['image'])
|
||||
response.choices[i].message.content[j]['image'] = f'data:image/jpg;base64,{b64_image}'
|
||||
|
||||
is_finished = all(response.choices[i].finish_reason for i in range(len(response.choices)))
|
||||
if 'stream' in response.__class__.__name__.lower():
|
||||
request_info['response'] += response.choices[0].delta.content
|
||||
else:
|
||||
request_info['response'] = response.choices[0].message.content
|
||||
if return_cmpl_response:
|
||||
response = response.to_cmpl_response()
|
||||
if is_finished:
|
||||
if args.log_interval > 0:
|
||||
self.infer_stats.update(response)
|
||||
if self.jsonl_writer:
|
||||
self.jsonl_writer.append(request_info)
|
||||
if self.args.verbose:
|
||||
logger.info(request_info)
|
||||
return response
|
||||
|
||||
def _set_request_config(self, request_config) -> None:
|
||||
default_request_config = self.args.get_request_config()
|
||||
if default_request_config is None:
|
||||
return
|
||||
for key, val in asdict(request_config).items():
|
||||
default_val = getattr(default_request_config, key)
|
||||
if default_val is not None and (val is None or isinstance(val, (list, tuple)) and len(val) == 0):
|
||||
setattr(request_config, key, default_val)
|
||||
|
||||
async def create_chat_completion(self,
|
||||
request: ChatCompletionRequest,
|
||||
raw_request: Request,
|
||||
*,
|
||||
return_cmpl_response: bool = False):
|
||||
args = self.args
|
||||
error_msg = (await self._check_model(request) or self._check_api_key(raw_request)
|
||||
or self._check_max_logprobs(request))
|
||||
if error_msg:
|
||||
return self.create_error_response(HTTPStatus.BAD_REQUEST, error_msg)
|
||||
infer_kwargs = self.infer_kwargs.copy()
|
||||
adapter_path = args.adapter_mapping.get(request.model)
|
||||
if adapter_path:
|
||||
infer_kwargs['adapter_request'] = AdapterRequest(request.model, adapter_path)
|
||||
|
||||
infer_request, request_config = request.parse()
|
||||
self._set_request_config(request_config)
|
||||
request_info = {'response': '', 'infer_request': infer_request.to_printable()}
|
||||
|
||||
def pre_infer_hook(kwargs):
|
||||
request_info['generation_config'] = kwargs['generation_config']
|
||||
return kwargs
|
||||
|
||||
infer_kwargs['pre_infer_hook'] = pre_infer_hook
|
||||
try:
|
||||
res_or_gen = await self.infer_async(infer_request, request_config, **infer_kwargs)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
logger.info(traceback.format_exc())
|
||||
return self.create_error_response(HTTPStatus.BAD_REQUEST, str(e))
|
||||
if request_config.stream:
|
||||
|
||||
async def _gen_wrapper():
|
||||
async for res in res_or_gen:
|
||||
res = self._post_process(request_info, res, return_cmpl_response)
|
||||
yield f'data: {json.dumps(asdict(res), ensure_ascii=False)}\n\n'
|
||||
yield 'data: [DONE]\n\n'
|
||||
|
||||
return StreamingResponse(_gen_wrapper(), media_type='text/event-stream')
|
||||
elif hasattr(res_or_gen, 'choices'):
|
||||
# instance of ChatCompletionResponse
|
||||
return self._post_process(request_info, res_or_gen, return_cmpl_response)
|
||||
else:
|
||||
return res_or_gen
|
||||
|
||||
async def create_completion(self, request: CompletionRequest, raw_request: Request):
|
||||
chat_request = ChatCompletionRequest.from_cmpl_request(request)
|
||||
return await self.create_chat_completion(chat_request, raw_request, return_cmpl_response=True)
|
||||
|
||||
async def create_embedding(self, request: EmbeddingRequest, raw_request: Request):
|
||||
chat_request = ChatCompletionRequest.from_cmpl_request(request)
|
||||
return await self.create_chat_completion(chat_request, raw_request, return_cmpl_response=True)
|
||||
|
||||
async def infer_handler(self, raw_request: Request):
|
||||
body = await raw_request.json()
|
||||
infer_requests = [RolloutInferRequest(**r) for r in body.get('infer_requests', [])]
|
||||
rc_data = body.get('request_config')
|
||||
request_config = RequestConfig(**rc_data) if rc_data else RequestConfig()
|
||||
results = await asyncio.gather(*[self.infer_async(req, request_config) for req in infer_requests])
|
||||
return results
|
||||
|
||||
def run(self):
|
||||
args = self.args
|
||||
self.jsonl_writer = JsonlWriter(args.result_path) if args.result_path else None
|
||||
logger.info(f'model_list: {self._get_model_list()}')
|
||||
uvicorn.run(
|
||||
self.app,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
ssl_keyfile=args.ssl_keyfile,
|
||||
ssl_certfile=args.ssl_certfile,
|
||||
log_level=args.log_level)
|
||||
|
||||
|
||||
def deploy_main(args: Optional[Union[List[str], DeployArguments]] = None) -> None:
|
||||
SwiftDeploy(args).main()
|
||||
|
||||
|
||||
def is_accessible(port: int):
|
||||
infer_client = InferClient(port=port)
|
||||
try:
|
||||
infer_client.get_model_list()
|
||||
except ClientConnectorError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _deploy_main(args):
|
||||
args._import_external_plugins()
|
||||
return deploy_main(args)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def run_deploy(args: DeployArguments, return_url: bool = False):
|
||||
if isinstance(args, DeployArguments) and args.__class__.__name__ == 'DeployArguments':
|
||||
deploy_args = args
|
||||
else:
|
||||
args_dict = asdict(args)
|
||||
parameters = inspect.signature(DeployArguments).parameters
|
||||
for k in list(args_dict.keys()):
|
||||
if k not in parameters or args_dict[k] is None:
|
||||
args_dict.pop(k)
|
||||
deploy_args = DeployArguments(**args_dict)
|
||||
|
||||
mp = multiprocessing.get_context('spawn')
|
||||
process = mp.Process(target=_deploy_main, args=(deploy_args, ))
|
||||
process.start()
|
||||
try:
|
||||
while not is_accessible(deploy_args.port):
|
||||
time.sleep(1)
|
||||
yield f'http://127.0.0.1:{deploy_args.port}/v1' if return_url else deploy_args.port
|
||||
finally:
|
||||
process.terminate()
|
||||
logger.info('The deployment process has been terminated.')
|
||||
@@ -0,0 +1,312 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import numpy as np
|
||||
from datasets import Dataset as HfDataset
|
||||
from tqdm import tqdm
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from swift.arguments import InferArguments
|
||||
from swift.dataset import DatasetLoader, load_dataset, sample_dataset
|
||||
from swift.infer_engine import AdapterRequest, InferRequest, RequestConfig, TransformersEngine
|
||||
from swift.metrics import InferStats, MeanMetric, compute_rouge_bleu
|
||||
from swift.utils import JsonlWriter, get_dist_setting, get_logger, is_dist, is_master, read_from_jsonl
|
||||
from ..base import SwiftPipeline
|
||||
from ..export import merge_lora
|
||||
from ..utils import get_cached_dataset, prepare_model_template
|
||||
from .utils import InferCliState
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftInfer(SwiftPipeline):
|
||||
args_class = InferArguments
|
||||
args: args_class
|
||||
|
||||
def __init__(self, args: Optional[Union[List[str], InferArguments]] = None) -> None:
|
||||
super().__init__(args)
|
||||
args = self.args
|
||||
if args.merge_lora:
|
||||
merge_lora(args, device_map='cpu')
|
||||
self.infer_kwargs = {}
|
||||
if args.infer_backend == 'vllm' and args.adapters:
|
||||
self.infer_kwargs['adapter_request'] = AdapterRequest('_lora', args.adapters[0])
|
||||
|
||||
if args.infer_backend == 'transformers':
|
||||
model, self.template = prepare_model_template(args)
|
||||
self.infer_engine = TransformersEngine(model, template=self.template, max_batch_size=args.max_batch_size)
|
||||
logger.info(f'model: {self.infer_engine.model}')
|
||||
else:
|
||||
self.template = args.get_template()
|
||||
self.infer_engine = self.get_infer_engine(args, self.template)
|
||||
self.random_state = np.random.RandomState(args.data_seed)
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
try:
|
||||
return super().__getattr__(key)
|
||||
except AttributeError:
|
||||
if 'infer_engine' in self.__dict__:
|
||||
return getattr(self.infer_engine, key)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def get_infer_engine(args: InferArguments, template=None, **extra_kwargs):
|
||||
infer_backend = extra_kwargs.pop('infer_backend', None) or args.infer_backend
|
||||
engine_kwargs = extra_kwargs.pop('engine_kwargs', {})
|
||||
kwargs = {
|
||||
'model_id_or_path': args.model,
|
||||
'model_type': args.model_type,
|
||||
'revision': args.model_revision,
|
||||
'torch_dtype': args.torch_dtype,
|
||||
'template': template,
|
||||
}
|
||||
if infer_backend in {'transformers', 'vllm'}:
|
||||
kwargs['reranker_use_activation'] = args.reranker_use_activation
|
||||
if infer_backend == 'transformers':
|
||||
infer_engine_cls = TransformersEngine
|
||||
kwargs.update(args.get_model_kwargs())
|
||||
if hasattr(args, 'max_batch_size'):
|
||||
kwargs.update({'max_batch_size': args.max_batch_size})
|
||||
elif infer_backend == 'vllm':
|
||||
from swift.infer_engine import VllmEngine
|
||||
infer_engine_cls = VllmEngine
|
||||
kwargs.update(args.get_vllm_engine_kwargs())
|
||||
seed = args.seed
|
||||
if is_dist():
|
||||
# Ensure that different data-parallel processes have different seeds.
|
||||
seed += get_dist_setting()[0] // args.vllm_tensor_parallel_size
|
||||
kwargs['distributed_executor_backend'] = 'external_launcher'
|
||||
kwargs['seed'] = seed
|
||||
elif infer_backend == 'sglang':
|
||||
from swift.infer_engine import SglangEngine
|
||||
infer_engine_cls = SglangEngine
|
||||
kwargs.update(args.get_sglang_engine_kwargs())
|
||||
elif infer_backend == 'lmdeploy':
|
||||
from swift.infer_engine import LmdeployEngine
|
||||
infer_engine_cls = LmdeployEngine
|
||||
kwargs.update(args.get_lmdeploy_engine_kwargs())
|
||||
else:
|
||||
raise ValueError(f'Inference backend `{infer_backend}` is not supported. '
|
||||
'Please use one of: transformers, vllm, sglang, lmdeploy.')
|
||||
if engine_kwargs:
|
||||
kwargs['engine_kwargs'] = kwargs.get('engine_kwargs') or {}
|
||||
kwargs['engine_kwargs'].update(engine_kwargs)
|
||||
kwargs.update(extra_kwargs)
|
||||
return infer_engine_cls(**kwargs)
|
||||
|
||||
def run(self) -> List[Dict[str, Any]]:
|
||||
args = self.args
|
||||
self.jsonl_writer = JsonlWriter(args.result_path) if args.result_path else None
|
||||
if args.eval_human:
|
||||
result = self.infer_cli()
|
||||
else:
|
||||
result = self.infer_dataset()
|
||||
if args.result_path:
|
||||
logger.info(f'The inference results have been saved to result_path: `{args.result_path}`.')
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def parse_data_from_response(response):
|
||||
if hasattr(response, 'choices'):
|
||||
return response.choices[0].message.content
|
||||
elif hasattr(response, 'data'):
|
||||
emb = response.data[0].embedding
|
||||
shape = len(emb)
|
||||
sample = str(emb)
|
||||
if len(emb) > 6:
|
||||
sample = str(emb[:3])[:-1] + ', ..., ' + str(emb[-3:])[1:]
|
||||
return f'Embedding(shape: [1, {shape}]): {sample}'
|
||||
|
||||
def infer_single(self, infer_request: Union[InferRequest, Dict[str, Any]], request_config: RequestConfig) -> str:
|
||||
res_or_gen = self.infer([infer_request], request_config, use_tqdm=False, **self.infer_kwargs)[0]
|
||||
if request_config and request_config.stream:
|
||||
response = ''
|
||||
for res in res_or_gen:
|
||||
delta = res.choices[0].delta.content
|
||||
print(delta, end='', flush=True)
|
||||
response += delta
|
||||
print()
|
||||
else:
|
||||
response = self.parse_data_from_response(res_or_gen)
|
||||
print(response)
|
||||
print('-' * 50)
|
||||
return response
|
||||
|
||||
def infer_cli(self) -> List[Dict[str, Any]]:
|
||||
args = self.args
|
||||
template = self.template
|
||||
request_config = args.get_request_config()
|
||||
logger.info(f'request_config: {request_config}')
|
||||
|
||||
logger.info('Input `exit` or `quit` to exit the conversation.')
|
||||
logger.info('Input `multi-line` to switch to multi-line input mode.')
|
||||
logger.info('Input `reset-system` to reset the system and clear the history.')
|
||||
support_multi_round = template.template_meta.support_multi_round
|
||||
if support_multi_round:
|
||||
logger.info('Input `clear` to clear the history.')
|
||||
else:
|
||||
logger.info('The current template only supports single-round dialogues.')
|
||||
|
||||
infer_state = InferCliState()
|
||||
result_list = []
|
||||
while True:
|
||||
if not support_multi_round:
|
||||
infer_state.clear()
|
||||
query = infer_state.input_text()
|
||||
if query.strip().lower() in {'exit', 'quit'}:
|
||||
break
|
||||
query = infer_state.check_query(query)
|
||||
if query is None:
|
||||
continue
|
||||
infer_state.add_query(query)
|
||||
if args.model_meta.is_multimodal:
|
||||
infer_state.input_mm_data()
|
||||
if args.model_meta.is_reward or args.task_type == 'prm':
|
||||
# reward model
|
||||
response = infer_state.input_text()
|
||||
infer_state.add_response(response)
|
||||
data = infer_state.to_dict()
|
||||
response = self.infer_single(data, request_config)
|
||||
data = {'response': response, **data}
|
||||
else:
|
||||
data = infer_state.to_dict()
|
||||
response = self.infer_single(data, request_config)
|
||||
infer_state.add_response(response)
|
||||
data['messages'].append({'role': 'assistant', 'content': response})
|
||||
data = {'response': response, **data}
|
||||
result_list.append(data)
|
||||
if self.jsonl_writer:
|
||||
self.jsonl_writer.append(data)
|
||||
|
||||
return result_list
|
||||
|
||||
def _prepare_val_dataset(self) -> HfDataset:
|
||||
args = self.args
|
||||
dataset_kwargs = args.get_dataset_kwargs()
|
||||
if args.cached_dataset or args.cached_val_dataset:
|
||||
_, val_datasets = get_cached_dataset(self.args)
|
||||
else:
|
||||
val_datasets = []
|
||||
if len(args.val_dataset) > 0:
|
||||
dataset_kwargs.pop('interleave_prob', None)
|
||||
_, val_dataset = load_dataset(
|
||||
args.val_dataset, split_dataset_ratio=1.0, shuffle=args.val_dataset_shuffle, **dataset_kwargs)
|
||||
val_datasets.append(val_dataset)
|
||||
elif args.dataset:
|
||||
_, val_dataset = load_dataset(
|
||||
args.dataset,
|
||||
split_dataset_ratio=args.split_dataset_ratio,
|
||||
shuffle=args.dataset_shuffle,
|
||||
**dataset_kwargs)
|
||||
val_datasets.append(val_dataset)
|
||||
assert len(val_datasets) > 0
|
||||
val_dataset = DatasetLoader.concat_datasets(val_datasets)
|
||||
val_dataset = sample_dataset(val_dataset, args.val_dataset_sample, args.dataset_shuffle, self.random_state)
|
||||
return val_dataset
|
||||
|
||||
def _calc_metric(self):
|
||||
args = self.args
|
||||
if not is_master():
|
||||
return
|
||||
data_list = read_from_jsonl(self.jsonl_writer.fpath)
|
||||
preds, labels = [], []
|
||||
for data in data_list:
|
||||
preds.append(data['response'])
|
||||
labels.append(data['labels'])
|
||||
if args.metric == 'acc':
|
||||
mean_metric = MeanMetric()
|
||||
for pred, label in zip(preds, labels):
|
||||
mean_metric.update(pred == label)
|
||||
res = {'acc': mean_metric.compute()['value']}
|
||||
elif args.metric == 'rouge':
|
||||
res = compute_rouge_bleu(preds, labels)
|
||||
logger.info(res)
|
||||
|
||||
def infer_dataset(self) -> List[Dict[str, Any]]:
|
||||
args = self.args
|
||||
request_config = args.get_request_config()
|
||||
logger.info(f'request_config: {request_config}')
|
||||
|
||||
val_dataset = self._prepare_val_dataset()
|
||||
logger.info(f'val_dataset: {val_dataset}')
|
||||
|
||||
self.infer_kwargs['metrics'] = [InferStats()]
|
||||
if request_config and request_config.stream:
|
||||
result_list = []
|
||||
for data in val_dataset:
|
||||
labels = InferRequest.remove_response(data['messages'])
|
||||
query = data['messages'][-1]['content']
|
||||
print(f'[QUERY] {query}')
|
||||
if labels:
|
||||
print(f'[LABELS] {labels}')
|
||||
print('[RESPONSE] ', end='')
|
||||
response = self.infer_single(data, request_config)
|
||||
data['messages'].append({'role': 'assistant', 'content': response})
|
||||
data = {'response': response, 'labels': labels, **data}
|
||||
result_list.append(data)
|
||||
if self.jsonl_writer:
|
||||
self.jsonl_writer.append(data)
|
||||
metrics = self.infer_kwargs.pop('metrics')
|
||||
print(metrics[0].compute())
|
||||
else:
|
||||
if args.write_batch_size <= 0:
|
||||
args.write_batch_size = len(val_dataset)
|
||||
if args.write_batch_size < len(val_dataset) and args.result_path:
|
||||
logger.info(f'args.result_path: {args.result_path}')
|
||||
prog_bar = tqdm(
|
||||
total=len(val_dataset), dynamic_ncols=True, disable=args.write_batch_size >= len(val_dataset))
|
||||
result_list = []
|
||||
idx = 0
|
||||
while idx < len(val_dataset):
|
||||
shard_size = min(args.write_batch_size, len(val_dataset) - idx)
|
||||
shard_dataset = val_dataset.select(range(idx, idx + shard_size))
|
||||
result = self._batch_infer(shard_dataset, request_config)
|
||||
if self.jsonl_writer:
|
||||
self.jsonl_writer.append(result, gather_obj=True)
|
||||
result_list += result
|
||||
idx += shard_size
|
||||
prog_bar.update(shard_size)
|
||||
prog_bar.close()
|
||||
metrics = self.infer_kwargs.pop('metrics')
|
||||
if result_list:
|
||||
metric = metrics[0].compute()
|
||||
print(f'[rank{args.rank}] {metric}' if args.rank >= 0 else str(metric))
|
||||
if args.metric is not None:
|
||||
self._calc_metric()
|
||||
return result_list
|
||||
|
||||
def _batch_infer(self, val_dataset, request_config):
|
||||
args = self.args
|
||||
result_list = []
|
||||
if args.infer_backend == 'vllm':
|
||||
rank = args.rank // args.vllm_tensor_parallel_size if args.rank >= 0 else -1
|
||||
data_parallel_size = args.global_world_size // args.vllm_tensor_parallel_size
|
||||
else:
|
||||
rank, data_parallel_size = args.rank, args.global_world_size
|
||||
# The dataset is insufficient for DP partitioning
|
||||
if len(val_dataset) < data_parallel_size:
|
||||
if rank >= len(val_dataset):
|
||||
return []
|
||||
data_parallel_size = len(val_dataset)
|
||||
if rank >= 0 and data_parallel_size > 1:
|
||||
val_dataset = val_dataset.shard(data_parallel_size, rank, contiguous=True)
|
||||
val_dataset = list(val_dataset)
|
||||
labels_list = []
|
||||
for data in val_dataset:
|
||||
if args.task_type == 'causal_lm':
|
||||
labels = InferRequest.remove_response(data['messages'])
|
||||
else:
|
||||
labels = data.pop('label', None)
|
||||
labels_list.append(labels)
|
||||
|
||||
resp_list = self.infer(val_dataset, request_config, use_tqdm=True, **self.infer_kwargs)
|
||||
if not (args.infer_backend == 'vllm' and rank >= 0
|
||||
and args.rank % args.vllm_tensor_parallel_size != 0): # DP & TP
|
||||
for data, resp, labels in zip(val_dataset, resp_list, labels_list):
|
||||
response = resp.choices[0].message.content
|
||||
data['messages'].append({'role': 'assistant', 'content': response})
|
||||
data = {'response': response, 'labels': labels, 'logprobs': resp.choices[0].logprobs, **data}
|
||||
result_list.append(data)
|
||||
return result_list
|
||||
|
||||
|
||||
def infer_main(args: Optional[Union[List[str], InferArguments]] = None):
|
||||
return SwiftInfer(args).main()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,116 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import re
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Literal, Optional
|
||||
|
||||
from swift.template import Messages
|
||||
from swift.utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
@dataclass
|
||||
class InferCliState:
|
||||
# None: use default-system. '': not use system.
|
||||
system: Optional[str] = None
|
||||
messages: Messages = field(default_factory=list) # not including system
|
||||
|
||||
images: List[str] = field(default_factory=list)
|
||||
audios: List[str] = field(default_factory=list)
|
||||
videos: List[str] = field(default_factory=list)
|
||||
|
||||
multiline_mode: bool = False
|
||||
input_system: bool = False
|
||||
|
||||
def clear(self):
|
||||
self.messages = []
|
||||
self.images = []
|
||||
self.audios = []
|
||||
self.videos = []
|
||||
|
||||
def add_query(self, query: str) -> None:
|
||||
role = 'user'
|
||||
if query.startswith('tool:'):
|
||||
role = 'tool'
|
||||
query = query[len('tool:'):]
|
||||
self.messages.append({'role': role, 'content': query})
|
||||
|
||||
def add_response(self, response: str) -> None:
|
||||
self.messages.append({'role': 'assistant', 'content': response})
|
||||
|
||||
def to_dict(self):
|
||||
infer_state = deepcopy(self)
|
||||
if infer_state.system is not None:
|
||||
infer_state.messages.insert(0, {'role': 'system', 'content': infer_state.system})
|
||||
return {
|
||||
'messages': infer_state.messages,
|
||||
'images': infer_state.images,
|
||||
'audios': infer_state.audios,
|
||||
'videos': infer_state.videos
|
||||
}
|
||||
|
||||
def input_mm_data(self) -> None:
|
||||
|
||||
def _input_mm_file(mm_type: Literal['image', 'video', 'audio']) -> str:
|
||||
a_an = 'an' if mm_type[0] in {'i', 'a'} else 'a'
|
||||
return input(f'Input {a_an} {mm_type} path or URL <<< ')
|
||||
|
||||
mm_types = ['image', 'video', 'audio']
|
||||
query = self.messages[-1]['content']
|
||||
mm_tags = re.findall('|'.join(f'<{mm_type}>' for mm_type in mm_types), query)
|
||||
# mm_tag -> mm_type/mm_key
|
||||
mm_mapping = {f'<{mm_type}>': (mm_type, f'{mm_type}s') for mm_type in mm_types}
|
||||
for mm_tag in mm_tags:
|
||||
mm_type, mm_key = mm_mapping[mm_tag]
|
||||
mm_val = getattr(self, mm_key)
|
||||
mm_val.append(_input_mm_file(mm_type))
|
||||
|
||||
@staticmethod
|
||||
def _input_multiline(prompt: str) -> str:
|
||||
query = ''
|
||||
stop_words = '#\n'
|
||||
while True:
|
||||
text = f'{input(prompt)}\n'
|
||||
prompt = ''
|
||||
if text.endswith(stop_words):
|
||||
query += text[:-len(stop_words)]
|
||||
break
|
||||
query += text
|
||||
return query
|
||||
|
||||
def input_text(self) -> str:
|
||||
if self.multiline_mode:
|
||||
addi_prompt = '[MS]' if self.input_system else '[M]'
|
||||
text = InferCliState._input_multiline(f'<<<{addi_prompt} ')
|
||||
else:
|
||||
addi_prompt = '[S]' if self.input_system else ''
|
||||
text = input(f'<<<{addi_prompt} ')
|
||||
return text
|
||||
|
||||
def check_query(self, query: str) -> Optional[str]:
|
||||
query_std = query.strip().lower()
|
||||
if self.input_system:
|
||||
if query == 'default-system':
|
||||
self.system = None
|
||||
else:
|
||||
self.system = query
|
||||
self.input_system = False
|
||||
query_std = 'clear'
|
||||
if query_std == 'clear':
|
||||
self.clear()
|
||||
return
|
||||
if query_std == '':
|
||||
return
|
||||
if query_std == 'reset-system':
|
||||
self.input_system = True
|
||||
return
|
||||
if query_std == 'multi-line':
|
||||
self.multiline_mode = True
|
||||
logger.info('End multi-line input with `#`.')
|
||||
logger.info('Input `single-line` to switch to single-line input mode.')
|
||||
return
|
||||
if query_std == 'single-line':
|
||||
self.multiline_mode = False
|
||||
return
|
||||
return query
|
||||
@@ -0,0 +1 @@
|
||||
from .sampling import sampling_main
|
||||
@@ -0,0 +1,59 @@
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from swift.arguments import SamplingArguments
|
||||
from swift.infer_engine import TransformersEngine
|
||||
from swift.ray_utils import RayHelper
|
||||
from swift.rewards import orms, prms
|
||||
from swift.utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class Sampler:
|
||||
|
||||
def __init__(self, input_args: SamplingArguments):
|
||||
self.args = input_args
|
||||
self.template = None
|
||||
self.processor = None
|
||||
self.prm_model = None
|
||||
self.orm_model = None
|
||||
self._prepare_model_tokenizer()
|
||||
self._prepare_template()
|
||||
self._prepare_prm()
|
||||
self._prepare_orm()
|
||||
|
||||
def _prepare_model_tokenizer(self):
|
||||
args = self.args
|
||||
_, self.processor = args.get_model_processor(load_model=False)
|
||||
|
||||
@RayHelper.function(group='prm')
|
||||
def _prepare_prm(self):
|
||||
if self.args.prm_model is None:
|
||||
self.prm_model = None
|
||||
logger.warning('prm_model is None.')
|
||||
elif self.args.prm_model in prms:
|
||||
self.prm_model = prms[self.args.prm_model]()
|
||||
else:
|
||||
self.prm_model = TransformersEngine(self.args.prm_model, max_batch_size=64)
|
||||
|
||||
@RayHelper.function(group='orm')
|
||||
def _prepare_orm(self):
|
||||
if self.args.orm_model is None:
|
||||
self.orm_model = None
|
||||
logger.warning('orm_model is None.')
|
||||
elif self.args.orm_model in orms:
|
||||
self.orm_model = orms[self.args.orm_model]()
|
||||
else:
|
||||
self.orm_model = TransformersEngine(self.args.orm_model, max_batch_size=64)
|
||||
|
||||
def _prepare_template(self) -> None:
|
||||
template = self.args.get_template(self.processor)
|
||||
self.template = template
|
||||
self.template.set_mode('train')
|
||||
|
||||
def truncate_input(self, slices: List[Dict[str, Any]]):
|
||||
"""Truncate the input rows to avoid hitting the max length of the policy model"""
|
||||
return slices
|
||||
|
||||
def do_sample(self, data):
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,160 @@
|
||||
import os
|
||||
from copy import deepcopy
|
||||
from openai import OpenAI
|
||||
from typing import List, Optional
|
||||
|
||||
from swift.infer_engine import InferRequest, RequestConfig
|
||||
from swift.ray_utils import RayHelper
|
||||
from .utils import get_messages_md5
|
||||
from .vanilla_sampler import VanillaSampler
|
||||
|
||||
|
||||
class OpenAIEngine:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: str,
|
||||
stream: bool = False,
|
||||
base_url: str = 'https://dashscope.aliyuncs.com/compatible-mode/v1',
|
||||
api_key: str = '',
|
||||
**kwargs,
|
||||
):
|
||||
self.model = model
|
||||
self.stream = stream
|
||||
self.client = OpenAI(api_key=api_key if api_key else os.getenv('OPENAI_API_KEY'), base_url=base_url, **kwargs)
|
||||
|
||||
def infer(
|
||||
self,
|
||||
infer_requests: List[InferRequest],
|
||||
request_config: Optional[RequestConfig] = None,
|
||||
):
|
||||
resp_contents = []
|
||||
for infer_request in infer_requests:
|
||||
completion = self.client.chat.completions.create(
|
||||
model=self.model,
|
||||
messages=infer_request['messages'],
|
||||
temperature=request_config.temperature,
|
||||
top_p=request_config.top_p,
|
||||
max_tokens=request_config.max_tokens,
|
||||
stream=self.stream,
|
||||
)
|
||||
reasoning_content = None
|
||||
if self.stream:
|
||||
reasoning_content = ''
|
||||
content = ''
|
||||
for chunk in completion:
|
||||
chunk_choices = chunk.choices
|
||||
if len(chunk_choices) == 0:
|
||||
continue
|
||||
reasoning_chunk = chunk_choices[0].delta.reasoning_content if hasattr(
|
||||
chunk_choices[0].delta, 'reasoning_content') else ''
|
||||
answer_chunk = chunk_choices[0].delta.content
|
||||
if reasoning_chunk:
|
||||
reasoning_content += reasoning_chunk
|
||||
elif answer_chunk:
|
||||
content += answer_chunk
|
||||
else:
|
||||
if hasattr(completion.choices[0].message, 'reasoning_content'):
|
||||
reasoning_content = completion.choices[0].message.reasoning_content
|
||||
content = completion.choices[0].message.content
|
||||
assert len(content) > 0, 'Empty completion'
|
||||
if reasoning_content:
|
||||
resp_content = f'<think>{reasoning_content}</think>\n\n<answer>{content}</answer>'
|
||||
else:
|
||||
resp_content = content
|
||||
resp_contents.append(resp_content)
|
||||
|
||||
return resp_contents
|
||||
|
||||
|
||||
@RayHelper.worker(group=['sampler', 'prm', 'orm'])
|
||||
class DistillSampler(VanillaSampler):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(VanillaSampler, self).__init__(*args, **kwargs)
|
||||
assert self.args.sampler_engine == 'client'
|
||||
self._prepare_sampler()
|
||||
self.caches = self.read_cache()
|
||||
|
||||
@RayHelper.function(group='sampler')
|
||||
def _prepare_sampler(self):
|
||||
self.infer_engine = OpenAIEngine(model=self.args.model, stream=self.args.stream, **self.args.engine_kwargs)
|
||||
self.infer_engine.strict = False
|
||||
|
||||
def _prepare_model_tokenizer(self):
|
||||
pass
|
||||
|
||||
def _prepare_template(self):
|
||||
pass
|
||||
|
||||
def extract_choice(self, resp):
|
||||
message = resp.choices[0].message
|
||||
if hasattr(message, 'reasoning_content'):
|
||||
reps_content = f'<think>{message.reasoning_content}</think>\n\n<answer>{message.content}</answer>'
|
||||
else:
|
||||
reps_content = message.content
|
||||
return reps_content
|
||||
|
||||
@RayHelper.function(
|
||||
group='sampler',
|
||||
dispatch=lambda n, i, data:
|
||||
([{
|
||||
'messages': data['messages'][i * len(data['messages']) // n:(i + 1) * len(data['messages']) // n]
|
||||
}], {}),
|
||||
collect='flatten')
|
||||
def generate(self, data):
|
||||
resp_all = []
|
||||
infer_requests = []
|
||||
sent = 0
|
||||
rows = self.convert_data_to_rows(data)
|
||||
for idx, row in enumerate(rows):
|
||||
row = deepcopy(row)
|
||||
messages = row['messages']
|
||||
uuid = get_messages_md5(row)
|
||||
if uuid in self.caches:
|
||||
choices = self.caches[uuid]['choices']
|
||||
if len(choices) == self.args.num_return_sequences:
|
||||
continue
|
||||
if self.args.system:
|
||||
if messages[0]['role'] == 'system':
|
||||
messages[0]['content'] = self.args.system
|
||||
else:
|
||||
messages.insert(0, {'role': 'system', 'content': self.args.system})
|
||||
if messages[-1]['role'] == 'assistant':
|
||||
messages = messages[:-1]
|
||||
|
||||
row['messages'] = messages
|
||||
infer_request = row
|
||||
for i in range(self.args.num_return_sequences):
|
||||
infer_requests.append(deepcopy(infer_request))
|
||||
sent += 1
|
||||
|
||||
request_config = RequestConfig(
|
||||
max_tokens=self.args.max_new_tokens,
|
||||
temperature=self.args.temperature,
|
||||
top_k=self.args.top_k,
|
||||
top_p=self.args.top_p,
|
||||
)
|
||||
|
||||
resp_list = []
|
||||
if len(infer_requests) > 0:
|
||||
resp_list = self.infer_engine.infer(infer_requests, request_config=request_config)
|
||||
|
||||
_cur = 0
|
||||
for idx, row in enumerate(rows):
|
||||
row = deepcopy(row)
|
||||
uuid = get_messages_md5(row)
|
||||
if uuid in self.caches:
|
||||
choices = self.caches[uuid]['choices']
|
||||
if len(choices) == self.args.num_return_sequences:
|
||||
row['choices'] = choices
|
||||
resp_all.append(row)
|
||||
continue
|
||||
|
||||
resps = row
|
||||
resps['choices'] = []
|
||||
for j in range(self.args.num_return_sequences * _cur, self.args.num_return_sequences * (_cur + 1)):
|
||||
resps['choices'].append(resp_list[j])
|
||||
resp_all.append(resps)
|
||||
_cur += 1
|
||||
return resp_all
|
||||
@@ -0,0 +1,104 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import time
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import SamplingArguments
|
||||
from swift.dataset import load_dataset
|
||||
from swift.utils import get_logger
|
||||
from ..base import SwiftPipeline
|
||||
from .distill_sampler import DistillSampler
|
||||
from .vanilla_sampler import VanillaSampler
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftSampling(SwiftPipeline):
|
||||
args_class = SamplingArguments
|
||||
args: args_class
|
||||
|
||||
def __init__(self, args: Optional[Union[List[str], SamplingArguments]] = None) -> None:
|
||||
super().__init__(args)
|
||||
self.args.save_args()
|
||||
os.makedirs(self.args.output_dir, exist_ok=True)
|
||||
self.cur_piece = 0
|
||||
self.total_piece = 1
|
||||
|
||||
if self.args.data_range:
|
||||
self.cur_piece, self.total_piece = self.args.data_range
|
||||
|
||||
if self.args.sampler_type == 'sample':
|
||||
self.sampler = VanillaSampler(self.args)
|
||||
elif self.args.sampler_type == 'distill':
|
||||
self.sampler = DistillSampler(self.args)
|
||||
else:
|
||||
raise ValueError(f'Unsupported sampler type: {self.args.sampler_type}')
|
||||
|
||||
def _get_dataset(self):
|
||||
args = self.args
|
||||
dataset_kwargs = args.get_dataset_kwargs()
|
||||
sampling_dataset, _ = load_dataset(
|
||||
args.dataset, split_dataset_ratio=0., shuffle=args.dataset_shuffle, **dataset_kwargs)
|
||||
logger.info(f'Sampling_dataset: {sampling_dataset}')
|
||||
dataset_len = len(sampling_dataset)
|
||||
piece_len = dataset_len // self.total_piece
|
||||
sampling_dataset = sampling_dataset.select(range(piece_len * self.cur_piece, piece_len * (self.cur_piece + 1)))
|
||||
return sampling_dataset
|
||||
|
||||
def run(self):
|
||||
os.makedirs(self.args.output_dir, exist_ok=True)
|
||||
iter_file = os.path.join(self.args.output_dir, self.args.output_file)
|
||||
resume_file = os.path.join(self.args.output_dir, self.args.output_file + '.resume')
|
||||
tmp_file = os.path.join(self.args.output_dir, self.args.output_file + '.tmp')
|
||||
ckpt_state_file = os.path.join(self.args.output_dir, 'ckpt_state.json')
|
||||
if os.path.exists(iter_file) and not self.args.override_exist_file:
|
||||
return
|
||||
|
||||
index_resume = -1
|
||||
write_mode = 'w'
|
||||
if self.args.resume:
|
||||
write_mode = 'a'
|
||||
if os.path.exists(resume_file):
|
||||
shutil.copyfile(resume_file, tmp_file)
|
||||
|
||||
if os.path.exists(ckpt_state_file):
|
||||
with open(ckpt_state_file, 'r', encoding='utf-8') as ckpt_state:
|
||||
data = json.load(ckpt_state)
|
||||
index_resume = data.get('index', -1)
|
||||
logger.info(f'Loaded index_resume: {index_resume}')
|
||||
else:
|
||||
if os.path.exists(tmp_file):
|
||||
os.remove(tmp_file)
|
||||
|
||||
dataset = self._get_dataset()
|
||||
dataset_len = len(dataset)
|
||||
total_iters = int(dataset_len // self.args.num_sampling_batch_size)
|
||||
|
||||
if self.args.num_sampling_batches is None or self.args.num_sampling_batches > total_iters:
|
||||
self.args.num_sampling_batches = total_iters
|
||||
|
||||
with open(tmp_file, write_mode) as f:
|
||||
for _index in range(self.args.num_sampling_batches):
|
||||
if _index <= index_resume:
|
||||
continue
|
||||
logger.info(f' Sampling index:{_index}')
|
||||
slices = dataset[self.args.num_sampling_batch_size * _index:self.args.num_sampling_batch_size
|
||||
* (_index + 1)]
|
||||
slices = self.sampler.truncate_input(slices)
|
||||
generated = self.sampler.do_sample(slices)
|
||||
f.writelines(generated)
|
||||
f.flush()
|
||||
shutil.copy(tmp_file, resume_file)
|
||||
with open(ckpt_state_file, 'w') as ckpt_state:
|
||||
json.dump({'index': _index}, ckpt_state)
|
||||
|
||||
if os.path.exists(iter_file):
|
||||
shutil.move(iter_file, iter_file + '.' + str(int(time.time())))
|
||||
shutil.move(resume_file, iter_file)
|
||||
logger.info(f'Sample file {iter_file} generated.')
|
||||
|
||||
|
||||
def sampling_main(args: Optional[Union[List[str], SamplingArguments]] = None):
|
||||
return SwiftSampling(args).main()
|
||||
@@ -0,0 +1,80 @@
|
||||
import hashlib
|
||||
import inspect
|
||||
import json
|
||||
import numpy as np
|
||||
from copy import copy
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from swift.infer_engine import ChatCompletionResponse, InferEngine, InferRequest, RequestConfig
|
||||
from swift.utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def get_messages_md5(row: Dict[str, Any]):
|
||||
row = copy(row)
|
||||
row.pop('choices', None)
|
||||
serialized = json.dumps(row, sort_keys=True)
|
||||
return hashlib.md5(serialized.encode('utf-8')).hexdigest()
|
||||
|
||||
|
||||
def get_reward(model: Any,
|
||||
infer_requests: List[InferRequest],
|
||||
request_config: RequestConfig = None,
|
||||
ground_truths: List[str] = None,
|
||||
threshold: Optional[float] = None):
|
||||
"""Get reward from an RM model.
|
||||
|
||||
Args:
|
||||
model: The model instance or an RM evaluator
|
||||
infer_requests: Infer requests sent to the model
|
||||
request_config: Infer config
|
||||
ground_truths: The ground truth list
|
||||
threshold: An optional threshold to generate the mask
|
||||
|
||||
Returns:
|
||||
Tuple
|
||||
Index 0: The min-max normalized scores matched the infer_requests
|
||||
Index 1: The mask filtered by the threshold
|
||||
"""
|
||||
infer_func = model.infer if isinstance(model, InferEngine) else model.__call__
|
||||
parameters = inspect.signature(infer_func).parameters
|
||||
gt_param = {}
|
||||
if 'ground_truths' in parameters:
|
||||
gt_param = {'ground_truths': ground_truths}
|
||||
if isinstance(infer_requests[0], dict):
|
||||
infer_requests = [InferRequest(messages=req['messages']) for req in infer_requests]
|
||||
rewards = infer_func(infer_requests, request_config=request_config, **gt_param)
|
||||
if isinstance(rewards[0], ChatCompletionResponse):
|
||||
print('reward:', rewards[0].choices[0].message.content)
|
||||
if isinstance(rewards[0].choices[0].message.content, str):
|
||||
rewards = [float(r.choices[0].message.content.strip('[]')) for r in rewards]
|
||||
elif isinstance(rewards[0].choices[0].message.content, list):
|
||||
rewards = [float(min(r.choices[0].message.content)) for r in rewards]
|
||||
else:
|
||||
rewards = [float(r.choices[0].message.content) for r in rewards]
|
||||
arr = []
|
||||
for reward in rewards:
|
||||
if isinstance(reward, (list, tuple)):
|
||||
arr.append(min(reward))
|
||||
else:
|
||||
arr.append(float(reward))
|
||||
|
||||
_mask = np.array([True] * len(arr))
|
||||
if threshold is not None:
|
||||
# > not >=, orm caller passes 0, which will cause error
|
||||
_mask = np.array([a > threshold for a in arr])
|
||||
|
||||
def normalize(arr):
|
||||
min_val = np.min(arr)
|
||||
max_val = np.max(arr)
|
||||
if min_val == max_val:
|
||||
if min_val == 0:
|
||||
constant_value = 0.0
|
||||
else:
|
||||
constant_value = min(1.0, min_val)
|
||||
return np.full_like(arr, fill_value=constant_value, dtype=np.float64)
|
||||
normalized = (arr - min_val) / (max_val - min_val + 1e-5)
|
||||
return normalized
|
||||
|
||||
return normalize(arr), _mask
|
||||
@@ -0,0 +1,233 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import json
|
||||
import numpy as np
|
||||
import os
|
||||
from copy import deepcopy
|
||||
|
||||
from swift.infer_engine import RequestConfig, TransformersEngine
|
||||
from swift.ray_utils import RayHelper
|
||||
from swift.utils import get_logger
|
||||
from .base import Sampler
|
||||
from .utils import get_messages_md5, get_reward
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
@RayHelper.worker(group=['sampler', 'prm', 'orm'])
|
||||
class VanillaSampler(Sampler):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._prepare_sampler()
|
||||
self.caches = self.read_cache()
|
||||
|
||||
@RayHelper.function(group='sampler')
|
||||
def _prepare_sampler(self):
|
||||
if self.args.sampler_engine == 'transformers':
|
||||
_Engine = TransformersEngine
|
||||
elif self.args.sampler_engine == 'vllm':
|
||||
from swift.infer_engine import VllmEngine
|
||||
_Engine = VllmEngine
|
||||
elif self.args.sampler_engine == 'lmdeploy':
|
||||
from swift.infer_engine import LmdeployEngine
|
||||
_Engine = LmdeployEngine
|
||||
elif self.args.sampler_engine == 'no':
|
||||
_Engine = None
|
||||
else:
|
||||
raise ValueError(f'Cannot find engine name: {self.args.sampler_engine}')
|
||||
self.infer_engine = None
|
||||
if _Engine:
|
||||
self.infer_engine = _Engine(
|
||||
self.args.model, model_type=self.args.model_type, template=self.template, **self.args.engine_kwargs)
|
||||
self.infer_engine.strict = False
|
||||
|
||||
@RayHelper.function(group='sampler')
|
||||
def read_cache(self):
|
||||
cache_files = self.args.cache_files
|
||||
caches = {}
|
||||
for file in cache_files:
|
||||
if not os.path.exists(file):
|
||||
logger.warning(f'Cache file does not exist: {file}')
|
||||
continue
|
||||
with open(file, 'r', encoding='utf-8') as f:
|
||||
for line in f.readlines():
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
|
||||
content = json.loads(line)
|
||||
uuid = content['id']
|
||||
messages = content['messages']
|
||||
if uuid not in caches:
|
||||
caches[uuid] = {'choices': []}
|
||||
assert messages[-1]['role'] == 'assistant'
|
||||
caches[uuid]['choices'].append(messages[-1]['content'])
|
||||
return caches
|
||||
|
||||
@staticmethod
|
||||
def convert_data_to_rows(data):
|
||||
rows = []
|
||||
key = list(data.keys())[0]
|
||||
data_len = len(data[key])
|
||||
for idx in range(data_len):
|
||||
row = {key: data[key][idx] for key in data}
|
||||
if row.get('images') and 'bytes' in row['images'][0]:
|
||||
row['images'] = [img['path'] for img in row['images']]
|
||||
rows.append(row)
|
||||
VanillaSampler.check_row_valid(rows)
|
||||
return rows
|
||||
|
||||
@staticmethod
|
||||
def check_row_valid(rows):
|
||||
for row in rows:
|
||||
assert not row.get('images') or all([isinstance(img, str) and img for img in row['images']])
|
||||
assert not row.get('videos') or all([isinstance(video, str) and video for video in row['videos']])
|
||||
assert not row.get('audios') or all([isinstance(audio, str) and audio for audio in row['audios']])
|
||||
|
||||
@RayHelper.function(
|
||||
group='sampler',
|
||||
dispatch=lambda n, i, data:
|
||||
([{
|
||||
'messages': data['messages'][i * len(data['messages']) // n:(i + 1) * len(data['messages']) // n]
|
||||
}], {}),
|
||||
collect='flatten')
|
||||
def generate(self, data):
|
||||
resp_all = []
|
||||
infer_requests = []
|
||||
sent = 0
|
||||
rows = self.convert_data_to_rows(data)
|
||||
for idx, row in enumerate(rows):
|
||||
row = deepcopy(row)
|
||||
messages = row['messages']
|
||||
uuid = get_messages_md5(row)
|
||||
if uuid in self.caches:
|
||||
choices = self.caches[uuid]['choices']
|
||||
if len(choices) == self.args.num_return_sequences:
|
||||
continue
|
||||
if self.args.system:
|
||||
if messages[0]['role'] == 'system':
|
||||
messages[0]['content'] = self.args.system
|
||||
else:
|
||||
messages.insert(0, {'role': 'system', 'content': self.args.system})
|
||||
if messages[-1]['role'] == 'assistant':
|
||||
messages = messages[:-1]
|
||||
|
||||
row['messages'] = messages
|
||||
infer_request = row
|
||||
for i in range(self.args.num_return_sequences):
|
||||
infer_requests.append(deepcopy(infer_request))
|
||||
sent += 1
|
||||
|
||||
request_config = RequestConfig(
|
||||
max_tokens=self.args.max_new_tokens,
|
||||
temperature=self.args.temperature,
|
||||
top_k=self.args.top_k,
|
||||
top_p=self.args.top_p,
|
||||
)
|
||||
|
||||
resp_list = []
|
||||
if len(infer_requests) > 0:
|
||||
resp_list = self.infer_engine.infer(infer_requests, request_config=request_config)
|
||||
|
||||
_cur = 0
|
||||
for idx, row in enumerate(rows):
|
||||
row = deepcopy(row)
|
||||
uuid = get_messages_md5(row)
|
||||
if uuid in self.caches:
|
||||
choices = self.caches[uuid]['choices']
|
||||
if len(choices) == self.args.num_return_sequences:
|
||||
row['choices'] = choices
|
||||
resp_all.append(row)
|
||||
continue
|
||||
|
||||
resps = row
|
||||
resps['choices'] = []
|
||||
for j in range(self.args.num_return_sequences * _cur, self.args.num_return_sequences * (_cur + 1)):
|
||||
if not isinstance(resp_list[j], Exception):
|
||||
resps['choices'].append(resp_list[j].choices[0].message.content)
|
||||
if resps['choices']:
|
||||
resp_all.append(resps)
|
||||
_cur += 1
|
||||
return resp_all
|
||||
|
||||
@RayHelper.function(group='orm', dispatch='slice', collect='flatten')
|
||||
def get_orm_score(self, infer_requests, ground_truth):
|
||||
return get_reward(
|
||||
self.orm_model, infer_requests, ground_truths=[ground_truth] * len(infer_requests), threshold=0.0)
|
||||
|
||||
@RayHelper.function(group='prm', dispatch='slice', collect='flatten')
|
||||
def get_prm_score(self, infer_requests, ground_truth):
|
||||
return get_reward(
|
||||
self.prm_model,
|
||||
infer_requests,
|
||||
ground_truths=[ground_truth] * len(infer_requests),
|
||||
threshold=self.args.prm_threshold)
|
||||
|
||||
def do_sample(self, data):
|
||||
generated = []
|
||||
resp_all = self.generate(data)
|
||||
for i, resps in enumerate(resp_all):
|
||||
choices = resps['choices']
|
||||
messages = resps['messages']
|
||||
uuid = get_messages_md5(resps)
|
||||
assert messages[-1]['role'] == 'assistant'
|
||||
ground_truth = messages[-1]['content']
|
||||
|
||||
infer_requests = []
|
||||
for decoded in choices:
|
||||
_resps = deepcopy(resps)
|
||||
_resps['messages'][-1]['content'] = decoded
|
||||
infer_requests.append(_resps)
|
||||
|
||||
_resps = deepcopy(resps)
|
||||
_resps['messages'][-1]['content'] = ground_truth
|
||||
infer_requests.append(_resps)
|
||||
if self.args.orm_model is not None:
|
||||
orm_score, _orm_mask = self.get_orm_score(infer_requests, ground_truth)
|
||||
else:
|
||||
orm_score = np.array([1.0] * len(infer_requests))
|
||||
_orm_mask = np.array([True] * len(infer_requests))
|
||||
if self.args.prm_model is not None:
|
||||
prm_score, _prm_mask = self.get_prm_score(infer_requests, ground_truth)
|
||||
else:
|
||||
prm_score = np.array([1.0] * len(infer_requests))
|
||||
_prm_mask = np.array([True] * len(infer_requests))
|
||||
|
||||
_mask = _orm_mask & _prm_mask
|
||||
if not any(_mask):
|
||||
continue
|
||||
|
||||
choices.append(ground_truth)
|
||||
choices = np.array(choices)
|
||||
|
||||
if self.args.orm_model is None and self.args.prm_model is None:
|
||||
positives = choices[:-1]
|
||||
for positive in positives:
|
||||
_resps = deepcopy(resps)
|
||||
_resps.pop('choices', None)
|
||||
_resps['id'] = uuid
|
||||
_resps['messages'][-1]['content'] = str(positive)
|
||||
generated.append(json.dumps(_resps, ensure_ascii=False) + '\n')
|
||||
else:
|
||||
score = np.array(prm_score) + np.array(orm_score * 10)
|
||||
sorted_indices = np.argsort(score)[::-1]
|
||||
pos_indexes = sorted_indices[0:self.args.n_best_to_keep]
|
||||
neg_index = sorted_indices[-1]
|
||||
pos_indexes = [int(i) for i in pos_indexes if _mask[i] and i != neg_index]
|
||||
logger.info(
|
||||
f'orm:{orm_score}, prm:{prm_score}, positive index: {pos_indexes}, negative index: {neg_index}')
|
||||
if self.args.easy_query_threshold is not None and sum([score > 0 for score in orm_score]) - 1 >= int(
|
||||
self.args.num_return_sequences * self.args.easy_query_threshold):
|
||||
continue
|
||||
if len(pos_indexes) > 0:
|
||||
positives = choices[pos_indexes]
|
||||
negative = choices[neg_index]
|
||||
for positive in positives:
|
||||
_resps = deepcopy(resps)
|
||||
messages = deepcopy(messages)
|
||||
_resps.pop('choices', None)
|
||||
_resps['messages'][-1]['content'] = str(positive)
|
||||
_resps['rejected_response'] = str(negative)
|
||||
_resps['id'] = uuid
|
||||
generated.append(json.dumps(_resps, ensure_ascii=False) + '\n')
|
||||
return generated
|
||||
@@ -0,0 +1,5 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
from .kto import prepare_kto_dataset
|
||||
from .pretrain import SwiftPretrain, pretrain_main
|
||||
from .rlhf import SwiftRLHF, rlhf_main
|
||||
from .sft import SwiftSft, sft_main
|
||||
@@ -0,0 +1,78 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import warnings
|
||||
from datasets import Dataset as HfDataset
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from swift.dataset import RowPreprocessor
|
||||
from swift.utils import get_dist_setting, get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class KTOPreprocessor(RowPreprocessor):
|
||||
|
||||
def batched_preprocess(self, batched_row: Dict[str, Any], **kwargs) -> Dict[str, Any]:
|
||||
batched_row = dict(batched_row)
|
||||
messages = batched_row['messages']
|
||||
batch_size = len(messages)
|
||||
kl_messages = [messages[-1]] + messages[:-1]
|
||||
|
||||
kl_response = []
|
||||
for i in range(batch_size):
|
||||
kl_message = kl_messages[i][-1]
|
||||
assert kl_message['role'] == 'assistant'
|
||||
kl_response.append(kl_message['content'])
|
||||
# The name rejected_response is just for convenience in processing.
|
||||
batched_row['rejected_response'] = kl_response
|
||||
|
||||
return batched_row
|
||||
|
||||
|
||||
def _get_kl_dataset(dataset: Optional[HfDataset],
|
||||
total_batch_size: int,
|
||||
num_proc: int,
|
||||
seed: Optional[int] = None) -> Optional[HfDataset]:
|
||||
# Shift one position to the right in each batch.
|
||||
if dataset is None:
|
||||
return
|
||||
dataset = dataset.shuffle(seed)
|
||||
return KTOPreprocessor()(dataset, batch_size=total_batch_size, num_proc=num_proc)
|
||||
|
||||
|
||||
def prepare_kto_dataset(args, train_dataset, val_dataset):
|
||||
if args.loss_type != 'apo_zero_unpaired':
|
||||
world_size = get_dist_setting()[2]
|
||||
if hasattr(args, 'global_batch_size') and args.global_batch_size is not None:
|
||||
total_batch_size = args.global_batch_size
|
||||
else:
|
||||
total_batch_size = (world_size * args.per_device_train_batch_size * args.gradient_accumulation_steps)
|
||||
if total_batch_size <= 1:
|
||||
raise ValueError('Batch size is 1 (too small). KTO will not work properly because the KL term '
|
||||
'will be equivalent to the implied reward.')
|
||||
train_dataset = _get_kl_dataset(train_dataset, total_batch_size, args.dataset_num_proc, args.data_seed)
|
||||
val_dataset = _get_kl_dataset(val_dataset, total_batch_size, args.dataset_num_proc, args.data_seed)
|
||||
|
||||
label = train_dataset['label']
|
||||
num_desirable = max(sum(label), 1)
|
||||
num_undesirable = max(len(label) - num_desirable, 1) # "label" is binary
|
||||
|
||||
if num_desirable != num_undesirable:
|
||||
# The lower and upper bounds come from Eq. (8) of https://huggingface.co/papers/2402.01306
|
||||
des_weight_lower_bound = round((num_undesirable * args.undesirable_weight / num_desirable) * 1, 2)
|
||||
des_weight_upper_bound = round((num_undesirable * args.undesirable_weight / num_desirable) * 1.33, 2)
|
||||
und_weight_lower_bound = round((num_desirable * args.desirable_weight / num_undesirable) / 1.33, 2)
|
||||
und_weight_upper_bound = round((num_desirable * args.desirable_weight / num_undesirable) / 1, 2)
|
||||
|
||||
des_weight_in_range = des_weight_lower_bound <= args.desirable_weight <= des_weight_upper_bound
|
||||
und_weight_in_range = und_weight_lower_bound <= args.undesirable_weight <= und_weight_upper_bound
|
||||
|
||||
if not (des_weight_in_range or und_weight_in_range):
|
||||
logger.info(f'desirable_weight: {args.desirable_weight}, undesirable_weight: {args.undesirable_weight}')
|
||||
warnings.warn(
|
||||
f"""
|
||||
You have different amounts of desirable/positive and undesirable/negative examples but the
|
||||
weights on the desirable and undesirable losses don't seem to be in an ideal range. Based
|
||||
on your data, we recommend EITHER desirable_weight in [{des_weight_lower_bound}, {des_weight_upper_bound}]
|
||||
or undesirable_weight in [{und_weight_lower_bound}, {und_weight_upper_bound}] (but NOT BOTH).
|
||||
See the documentation on how to optimally set these weights.""", UserWarning)
|
||||
return train_dataset, val_dataset
|
||||
@@ -0,0 +1,17 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import PretrainArguments
|
||||
from swift.utils import get_logger
|
||||
from .sft import SwiftSft
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftPretrain(SwiftSft):
|
||||
args_class = PretrainArguments
|
||||
args: args_class
|
||||
|
||||
|
||||
def pretrain_main(args: Optional[Union[List[str], PretrainArguments]] = None):
|
||||
return SwiftPretrain(args).main()
|
||||
@@ -0,0 +1,248 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import os
|
||||
import peft
|
||||
from contextlib import nullcontext
|
||||
from packaging import version
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import BaseArguments, RLHFArguments
|
||||
from swift.dataset import DatasetLoader, load_dataset
|
||||
from swift.model import get_model_info_meta
|
||||
from swift.sequence_parallel import sequence_parallel
|
||||
from swift.tuner_plugin import Tuner, tuners_map
|
||||
from swift.tuners import Swift
|
||||
from swift.utils import (HfConfigFactory, disable_deepspeed_zero3, get_logger, get_model_parameter_info,
|
||||
safe_snapshot_download)
|
||||
from ..utils import prepare_adapter
|
||||
from .kto import prepare_kto_dataset
|
||||
from .sft import SwiftSft
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
class SwiftRLHF(SwiftSft):
|
||||
args_class = RLHFArguments
|
||||
args: args_class
|
||||
|
||||
@staticmethod
|
||||
def _get_model_task_type(model_dir):
|
||||
task_type = None
|
||||
num_labels = None
|
||||
if os.path.exists(os.path.join(model_dir, 'args.json')):
|
||||
model_args = BaseArguments.from_pretrained(model_dir)
|
||||
if hasattr(model_args, 'task_type'):
|
||||
task_type = model_args.task_type
|
||||
if hasattr(model_args, 'num_labels'):
|
||||
num_labels = model_args.num_labels
|
||||
if task_type == 'seq_cls' and num_labels is None:
|
||||
num_labels = 1
|
||||
else:
|
||||
from transformers import AutoConfig
|
||||
model_config = AutoConfig.from_pretrained(model_dir, trust_remote_code=True)
|
||||
if hasattr(model_config, 'architectures') and model_config.architectures:
|
||||
if any('sequenceclassification' in arch.lower() for arch in model_config.architectures):
|
||||
task_type = 'seq_cls'
|
||||
num_labels = getattr(model_config, 'num_labels', None) or 1
|
||||
|
||||
if task_type is None:
|
||||
if hasattr(model_config, 'num_labels'):
|
||||
num_labels = model_config.num_labels
|
||||
# PretrainedConfig default num_labels = 2
|
||||
if num_labels == 1:
|
||||
task_type = 'seq_cls'
|
||||
return task_type, num_labels
|
||||
|
||||
def _prepare_single_model(self, key, origin_key, model_type, model_revision):
|
||||
args = self.args
|
||||
origin_key = origin_key or key
|
||||
model_id_or_path = getattr(args, f'{key}_model')
|
||||
if model_id_or_path is None:
|
||||
return
|
||||
|
||||
if args.rlhf_type == 'ppo' and key == 'reward' and isinstance(model_id_or_path, (list, tuple)):
|
||||
assert len(model_id_or_path) == 1, f'model_id_or_path: {model_id_or_path}'
|
||||
model_id_or_path = model_id_or_path[0]
|
||||
|
||||
if model_type is None:
|
||||
model_info, _ = get_model_info_meta(model_id_or_path)
|
||||
model_type = model_info.model_type
|
||||
|
||||
if isinstance(model_id_or_path, list):
|
||||
# value model in PPO
|
||||
model_id_or_path = model_id_or_path[0]
|
||||
|
||||
model_dir = safe_snapshot_download(
|
||||
model_id_or_path=model_id_or_path,
|
||||
revision=model_revision,
|
||||
download_model=False,
|
||||
use_hf=args.use_hf,
|
||||
hub_token=args.hub_token,
|
||||
)
|
||||
task_type, num_labels = self._get_model_task_type(model_dir)
|
||||
context = nullcontext()
|
||||
if key == 'teacher' and args.teacher_deepspeed:
|
||||
if args.teacher_deepspeed.get('zero_optimization', {}).get('stage') != 3:
|
||||
context = disable_deepspeed_zero3()
|
||||
with context:
|
||||
model, processor = args.get_model_processor(
|
||||
model=model_id_or_path,
|
||||
model_type=model_type,
|
||||
revision=model_revision,
|
||||
task_type=task_type,
|
||||
num_labels=num_labels)
|
||||
|
||||
adapters = args.adapters if key == 'ref' else args.reward_adapters
|
||||
model = prepare_adapter(args, model, adapters)
|
||||
if origin_key in {'ref', 'reward', 'teacher'}:
|
||||
if self.args.sequence_parallel_size > 1:
|
||||
sequence_parallel.prepare(
|
||||
self.args.sequence_parallel_size, model, processor, padding_free=args.padding_free)
|
||||
model.requires_grad_(False).eval()
|
||||
else:
|
||||
model = self.prepare_model(args, model, task_type=task_type)
|
||||
logger.info(f'value_model: {model}')
|
||||
model_parameter_info = get_model_parameter_info(model)
|
||||
self.train_msg['value_model_parameter_info'] = model_parameter_info
|
||||
logger.info(f'value_model_parameter_info: {model_parameter_info}')
|
||||
|
||||
HfConfigFactory.set_config_attr(model.config, 'use_cache', False)
|
||||
return model, processor
|
||||
|
||||
def _prepare_model_tokenizer(self):
|
||||
# prepare ref/reward/value model
|
||||
args = self.args
|
||||
# Handle ref and value models
|
||||
for key in ['ref', 'value', 'teacher']:
|
||||
setattr(self, f'{key}_model', None)
|
||||
if key == 'ref' and args.rlhf_type == 'gkd':
|
||||
continue
|
||||
if key == 'value' and args.rlhf_type != 'ppo':
|
||||
continue
|
||||
if key == 'teacher' and args.rlhf_type not in ['gkd', 'grpo']:
|
||||
continue
|
||||
model_key = 'reward' if key == 'value' else key
|
||||
model_type = getattr(args, f'{model_key}_model_type')
|
||||
model_revision = getattr(args, f'{model_key}_model_revision')
|
||||
if key == 'value':
|
||||
model_type = model_type[0] if model_type else None
|
||||
model_revision = model_revision[0] if model_revision else None
|
||||
|
||||
result = self._prepare_single_model(model_key, key, model_type, model_revision)
|
||||
if result is not None:
|
||||
model, _ = result
|
||||
setattr(self, f'{key}_model', model)
|
||||
|
||||
# Handle reward model(s)
|
||||
self.reward_model = None
|
||||
if hasattr(args, 'reward_model') and args.reward_model is not None:
|
||||
rms = args.reward_model if isinstance(args.reward_model, list) else [args.reward_model]
|
||||
num_rms = len(rms)
|
||||
rm_types = args.reward_model_type if args.reward_model_type else [None] * num_rms
|
||||
rm_templates = args.reward_template if args.reward_template else [None] * num_rms
|
||||
rm_revisions = args.reward_model_revision if args.reward_model_revision else [None] * num_rms
|
||||
assert len(rms) == len(rm_types) == len(rm_templates) == len(rm_revisions)
|
||||
|
||||
self.reward_model = []
|
||||
if args.rlhf_type == 'grpo':
|
||||
self.reward_template = []
|
||||
|
||||
for reward_model_path, rm_type, rm_template, rm_revision in zip(rms, rm_types, rm_templates, rm_revisions):
|
||||
args.reward_model = reward_model_path # Temporarily set for prepare_single_model
|
||||
result = self._prepare_single_model('reward', None, rm_type, rm_revision)
|
||||
if result is not None:
|
||||
model, processor = result
|
||||
self.reward_model.append(model)
|
||||
|
||||
if args.rlhf_type == 'grpo':
|
||||
template_type = rm_template or processor.model_meta.template
|
||||
reward_template = self.args.get_template(processor, template_type=template_type)
|
||||
if reward_template.use_model:
|
||||
reward_template.model = model
|
||||
self.reward_template.append(reward_template)
|
||||
args.reward_model = rms # Restore original value
|
||||
if args.rlhf_type != 'grpo' and self.reward_model:
|
||||
assert len(self.reward_model) <= 1
|
||||
self.reward_model = self.reward_model[0]
|
||||
|
||||
super()._prepare_model_tokenizer()
|
||||
|
||||
@classmethod
|
||||
def prepare_model(cls, args, model, *, template=None, train_dataset=None, task_type=None):
|
||||
model = super().prepare_model(args, model, template=template, train_dataset=train_dataset, task_type=task_type)
|
||||
if args.ref_adapters:
|
||||
if args.tuner_type in tuners_map:
|
||||
tuner: Tuner = tuners_map[args.tuner_type]
|
||||
else:
|
||||
tuner = Swift
|
||||
assert len(args.ref_adapters) == 1, f'args.ref_adapters: {args.ref_adapters}'
|
||||
# is_trainable: fix peft0.18.1
|
||||
kwargs = {}
|
||||
if version.parse(peft.__version__) >= version.parse('0.18'):
|
||||
kwargs['is_trainable'] = True
|
||||
model = tuner.from_pretrained(model, args.ref_adapters[0], adapter_name='ref_adapter', **kwargs)
|
||||
assert args.rlhf_type in {'dpo', 'kto',
|
||||
'grpo'}, 'Currently, only DPO, KTO, and GRPO support `ref_adapters`.'
|
||||
args.training_args.ref_adapter_name = 'ref_adapter'
|
||||
return model
|
||||
|
||||
def _prepare_template(self) -> None:
|
||||
args = self.args
|
||||
super()._prepare_template()
|
||||
mode_mapping = {'kto': 'kto', 'gkd': 'train', 'ppo': 'transformers', 'grpo': 'train'}
|
||||
self.template.set_mode(mode_mapping.get(args.rlhf_type, 'rlhf'))
|
||||
|
||||
if args.rlhf_type == 'ppo':
|
||||
args.training_args.stop_token_id = self.template.template_meta.stop_token_id
|
||||
|
||||
def _get_dataset(self):
|
||||
args = self.args
|
||||
train_dataset, val_dataset = super()._get_dataset()
|
||||
if args.rlhf_type == 'kto':
|
||||
train_dataset, val_dataset = prepare_kto_dataset(args, train_dataset, val_dataset)
|
||||
return train_dataset, val_dataset
|
||||
|
||||
def _prepare_chord_sft_dataset(self):
|
||||
# prepare expert sft dataset for chord
|
||||
args = self.args
|
||||
assert hasattr(args, 'chord_sft_dataset') and args.chord_sft_dataset
|
||||
dataset_kwargs = args.get_dataset_kwargs()
|
||||
chord_sft_datasets = []
|
||||
# TODO: validatition
|
||||
chord_sft_dataset, _ = load_dataset(
|
||||
args.chord_sft_dataset, split_dataset_ratio=0, shuffle=args.dataset_shuffle, **dataset_kwargs)
|
||||
chord_sft_dataset, _ = self._encode_dataset(chord_sft_dataset, None, pre_process=True)
|
||||
chord_sft_datasets.append(chord_sft_dataset)
|
||||
chord_sft_dataset = DatasetLoader.concat_datasets(chord_sft_datasets)
|
||||
datasets = [chord_sft_dataset, None]
|
||||
datasets = self._post_process_datasets(datasets)
|
||||
return datasets
|
||||
|
||||
def _get_trainer_kwargs(self):
|
||||
trainer_kwargs = {}
|
||||
for key in ['ref', 'reward', 'value', 'teacher']:
|
||||
key = f'{key}_model'
|
||||
model = getattr(self, key, None)
|
||||
if model or self.args.rlhf_type == 'ppo' and key != 'teacher_model':
|
||||
trainer_kwargs[key] = model
|
||||
if hasattr(self, 'reward_template'):
|
||||
trainer_kwargs['reward_template'] = self.reward_template
|
||||
if self.args.rlhf_type in ['grpo', 'gkd']:
|
||||
trainer_kwargs['vllm_client'] = self.args.vllm_client
|
||||
if self.args.rlhf_type == 'grpo':
|
||||
trainer_kwargs['reward_funcs'] = self.args.reward_funcs
|
||||
if self.args.chord_sft_dataset:
|
||||
trainer_kwargs['chord_sft_dataset'], _ = self._prepare_chord_sft_dataset()
|
||||
# Teacher wiring shared by GKD and GRPO+OPD-RL (gkd_logits_topk is GKD-only).
|
||||
if self.args.rlhf_type in ['gkd', 'grpo']:
|
||||
if self.args.teacher_deepspeed:
|
||||
trainer_kwargs['teacher_deepspeed_config'] = self.args.teacher_deepspeed
|
||||
if self.args.teacher_model_server:
|
||||
trainer_kwargs['teacher_model_server'] = self.args.teacher_model_server
|
||||
trainer_kwargs['teacher_use_disable_adapter'] = getattr(self.args, '_teacher_use_disable_adapter', False)
|
||||
if self.args.rlhf_type == 'gkd':
|
||||
trainer_kwargs['gkd_logits_topk'] = self.args.gkd_logits_topk
|
||||
return trainer_kwargs
|
||||
|
||||
|
||||
def rlhf_main(args: Optional[Union[List[str], RLHFArguments]] = None):
|
||||
return SwiftRLHF(args).main()
|
||||
@@ -0,0 +1,341 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import os
|
||||
from datasets import Dataset as HfDataset
|
||||
from typing import List, Optional, Union
|
||||
|
||||
from swift.arguments import SftArguments
|
||||
from swift.dataset import (AddLengthPreprocessor, DatasetLoader, EncodePreprocessor, IterablePackingDataset,
|
||||
LazyLLMDataset, PackingDataset)
|
||||
from swift.infer_engine import prepare_generation_config
|
||||
from swift.ray_utils import RayHelper
|
||||
from swift.sequence_parallel import sequence_parallel
|
||||
from swift.trainers import TrainerFactory
|
||||
from swift.utils import append_to_jsonl, get_logger, get_model_parameter_info, is_master, plot_images, stat_array
|
||||
from ..base import SwiftPipeline
|
||||
from ..utils import get_cached_dataset
|
||||
from .tuner import TunerMixin
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
@RayHelper.worker(group=['default'])
|
||||
class SwiftSft(SwiftPipeline, TunerMixin):
|
||||
args_class = SftArguments
|
||||
args: args_class
|
||||
|
||||
def __init__(self, args: Optional[Union[List[str], SftArguments]] = None) -> None:
|
||||
super().__init__(args)
|
||||
self.train_msg = {}
|
||||
self._prepare_model_tokenizer()
|
||||
self._prepare_template()
|
||||
self._prepare_flash_ckpt()
|
||||
|
||||
@RayHelper.function(group='default')
|
||||
def _prepare_flash_ckpt(self):
|
||||
if self.args.use_flash_ckpt:
|
||||
try:
|
||||
import dlrover.trainer.torch.flash_checkpoint.hf_trainer
|
||||
except ImportError:
|
||||
raise ValueError('Please install dlrover to use flash ckpt `pip install dlrover[k8s,torch]')
|
||||
|
||||
def _prepare_generation_config(self):
|
||||
args = self.args
|
||||
self.model.origin_generation_config = self.model.generation_config
|
||||
self.model.generation_config = prepare_generation_config(self.model.generation_config,
|
||||
args.get_request_config(), self.tokenizer)
|
||||
logger.info(f'model.generation_config: {self.model.generation_config}')
|
||||
|
||||
@RayHelper.function(group='default')
|
||||
def _prepare_model_tokenizer(self, **kwargs):
|
||||
args = self.args
|
||||
self.model, self.processor = args.get_model_processor(**kwargs)
|
||||
if args.sequence_parallel_size > 1:
|
||||
sequence_parallel.prepare(
|
||||
args.sequence_parallel_size, model=self.model, tokenizer=self.processor, padding_free=args.padding_free)
|
||||
if self.model is None:
|
||||
return
|
||||
if hasattr(self.model, 'hf_device_map'):
|
||||
logger.info(f'model.hf_device_map: {self.model.hf_device_map}')
|
||||
|
||||
logger.info(f'model_info: {self.model.model_info}')
|
||||
|
||||
self._prepare_generation_config()
|
||||
|
||||
@RayHelper.function(group='default')
|
||||
def _prepare_template(self) -> None:
|
||||
args = self.args
|
||||
template = args.get_template(self.processor)
|
||||
template.set_mode('train')
|
||||
if template.use_model:
|
||||
template.model = self.model
|
||||
support_padding_free = template.support_padding_free
|
||||
if support_padding_free is None:
|
||||
support_padding_free = not args.model_meta.is_multimodal
|
||||
if (args.padding_free or args.packing) and not support_padding_free:
|
||||
raise ValueError(f'Template `{args.template}` does not support padding free or packing.')
|
||||
self.template = template
|
||||
|
||||
def _get_dataset(self):
|
||||
# The random shuffling of the training set occurs in the dataloader of the trainer.
|
||||
args = self.args
|
||||
train_dataset, val_dataset = args.load_dataset()
|
||||
if args.truncation_strategy == 'split':
|
||||
logger.info(f'train_dataset: {train_dataset}')
|
||||
logger.info(f'val_dataset: {val_dataset}')
|
||||
return train_dataset, val_dataset
|
||||
|
||||
def _save_val_dataset(self, val_dataset):
|
||||
args = self.args
|
||||
output_dir = getattr(args, 'output_dir', None) or getattr(args, 'save')
|
||||
if is_master() and isinstance(val_dataset, HfDataset) and not args.val_dataset:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
val_dataset_path = os.path.join(output_dir, 'val_dataset.jsonl')
|
||||
append_to_jsonl(val_dataset_path, val_dataset.to_list())
|
||||
logger.info(f'The split dataset from the training set will be saved at: `{val_dataset_path}`.')
|
||||
|
||||
@RayHelper.function(group='default')
|
||||
def _prepare_dataset(self):
|
||||
args = self.args
|
||||
# Defer encoding to the training phase
|
||||
pre_process = not (hasattr(args, 'rlhf_type') and args.rlhf_type in ['grpo', 'gkd'])
|
||||
if args.cached_dataset or args.cached_val_dataset:
|
||||
assert not args.streaming, 'Cached dataset does not support streaming.'
|
||||
train_datasets, val_datasets = get_cached_dataset(self.args)
|
||||
else:
|
||||
train_datasets, val_datasets = [], []
|
||||
if args.dataset or args.val_dataset:
|
||||
train_dataset, val_dataset = self._get_dataset()
|
||||
train_dataset, val_dataset = self._encode_dataset(train_dataset, val_dataset, pre_process=pre_process)
|
||||
if train_dataset is not None:
|
||||
train_datasets.append(train_dataset)
|
||||
if val_dataset is not None:
|
||||
val_datasets.append(val_dataset)
|
||||
train_dataset = DatasetLoader.concat_datasets(train_datasets)
|
||||
val_dataset = DatasetLoader.concat_datasets(val_datasets)
|
||||
if args.truncation_strategy != 'split':
|
||||
logger.info(f'train_dataset: {train_dataset}')
|
||||
logger.info(f'val_dataset: {val_dataset}')
|
||||
datasets = [train_dataset, val_dataset]
|
||||
if not pre_process:
|
||||
return datasets
|
||||
datasets = self._post_process_datasets(datasets)
|
||||
self._show_dataset(*datasets)
|
||||
return datasets
|
||||
|
||||
def _post_process_datasets(self, datasets: List) -> List:
|
||||
args = self.args
|
||||
predict_with_generate = getattr(args, 'predict_with_generate', False)
|
||||
|
||||
template = self.template
|
||||
for i, dataset in enumerate(datasets):
|
||||
if dataset is None:
|
||||
continue
|
||||
if i == 1 and predict_with_generate:
|
||||
# val_dataset
|
||||
continue
|
||||
if not args.streaming and args.truncation_strategy != 'split':
|
||||
dataset = LazyLLMDataset(dataset, template.encode, strict=args.strict, random_state=args.data_seed)
|
||||
if args.packing:
|
||||
packing_dataset_cls = IterablePackingDataset if args.streaming else PackingDataset
|
||||
dataset = packing_dataset_cls(
|
||||
template,
|
||||
dataset,
|
||||
num_proc=args.dataset_num_proc,
|
||||
packing_length=args.packing_length,
|
||||
packing_num_proc=args.packing_num_proc,
|
||||
packing_strategy=args.packing_strategy,
|
||||
strict=args.strict,
|
||||
load_from_cache_file=args.load_from_cache_file)
|
||||
elif args.streaming:
|
||||
preprocessor = EncodePreprocessor(template=template)
|
||||
dataset = preprocessor(
|
||||
dataset,
|
||||
num_proc=args.dataset_num_proc,
|
||||
load_from_cache_file=args.load_from_cache_file,
|
||||
strict=args.strict)
|
||||
datasets[i] = dataset
|
||||
return datasets
|
||||
|
||||
@RayHelper.function(group='default')
|
||||
def run(self):
|
||||
args = self.args
|
||||
train_dataset, val_dataset = self._prepare_dataset()
|
||||
|
||||
if args.task_type == 'seq_cls':
|
||||
args.problem_type = args.problem_type or getattr(self.model.config, 'problem_type', None)
|
||||
logger.info(f'args.problem_type: {args.problem_type}')
|
||||
args.save_args()
|
||||
|
||||
# Some tuners require train_dataset and data_collator for preparation: LoRA-GA
|
||||
self.model = self.prepare_model(self.args, self.model, template=self.template, train_dataset=train_dataset)
|
||||
logger.info(f'model: {self.model}')
|
||||
model_parameter_info = get_model_parameter_info(self.model)
|
||||
self.train_msg['model_parameter_info'] = model_parameter_info
|
||||
logger.info(f'model_parameter_info: {model_parameter_info}')
|
||||
|
||||
trainer_cls = TrainerFactory.get_trainer_cls(args)
|
||||
trainer = trainer_cls(
|
||||
model=self.model,
|
||||
args=self.args.training_args,
|
||||
template=self.template,
|
||||
train_dataset=train_dataset,
|
||||
eval_dataset=val_dataset,
|
||||
**self._get_trainer_kwargs(),
|
||||
)
|
||||
return self.train(trainer)
|
||||
|
||||
def _get_trainer_kwargs(self):
|
||||
return {}
|
||||
|
||||
def _handle_trainer_state(self, trainer, is_write_rank: bool):
|
||||
state = trainer.state
|
||||
if hasattr(state, 'last_model_checkpoint'):
|
||||
if self.args.create_checkpoint_symlink:
|
||||
last_checkpoint = os.path.join(self.args.output_dir, 'last')
|
||||
best_checkpoint = os.path.join(self.args.output_dir, 'best')
|
||||
if is_write_rank:
|
||||
os.symlink(state.last_model_checkpoint, last_checkpoint)
|
||||
os.symlink(state.best_model_checkpoint, best_checkpoint)
|
||||
state.last_model_checkpoint = last_checkpoint
|
||||
state.best_model_checkpoint = best_checkpoint
|
||||
else:
|
||||
state.last_model_checkpoint = None
|
||||
logger.info_if(f'last_model_checkpoint: {state.last_model_checkpoint}', cond=is_write_rank)
|
||||
logger.info_if(f'best_model_checkpoint: {state.best_model_checkpoint}', cond=is_write_rank)
|
||||
|
||||
def _save_trainer_state(self, trainer):
|
||||
training_args = trainer.args
|
||||
state = trainer.state
|
||||
self._handle_trainer_state(trainer, is_master())
|
||||
|
||||
if is_master():
|
||||
# Visualization
|
||||
if 'tensorboard' in training_args.report_to:
|
||||
images_dir = os.path.join(training_args.output_dir, 'images')
|
||||
logger.info(f'images_dir: {images_dir}')
|
||||
plot_images(images_dir, training_args.logging_dir, ['train/loss'], 0.9)
|
||||
if training_args.push_to_hub:
|
||||
trainer.push_to_hub()
|
||||
|
||||
self.train_msg.update({
|
||||
'last_model_checkpoint': state.last_model_checkpoint,
|
||||
'best_model_checkpoint': state.best_model_checkpoint,
|
||||
'best_metric': state.best_metric,
|
||||
'global_step': state.global_step,
|
||||
'log_history': state.log_history,
|
||||
'memory': getattr(state, 'max_memory', None),
|
||||
})
|
||||
if is_master():
|
||||
jsonl_path = os.path.join(training_args.output_dir, 'logging.jsonl')
|
||||
append_to_jsonl(jsonl_path, self.train_msg, strict=False)
|
||||
return self.train_msg
|
||||
|
||||
def _get_resume_checkpoint(self, trainer):
|
||||
args = trainer.args
|
||||
if args.resume_from_checkpoint:
|
||||
return args.resume_from_checkpoint
|
||||
resume_checkpoint = None
|
||||
# If flash checkpoint is enabled, try to resume from the last complete checkpoint.
|
||||
# If the previous training finished, resume_checkpoint stays None.
|
||||
if args.use_flash_ckpt:
|
||||
# resume_checkpoint = <resume_dir>/checkpoint-<step>
|
||||
resume_checkpoint = trainer.get_resume_checkpoint()
|
||||
|
||||
# Elastic runs require a universal checkpoint; fall back when missing or incomplete.
|
||||
callbacks = set(getattr(args, 'callbacks', []))
|
||||
elastic_enabled = 'deepspeed_elastic' in callbacks
|
||||
if elastic_enabled and (resume_checkpoint is None
|
||||
or not os.path.exists(os.path.join(resume_checkpoint, 'latest_universal'))):
|
||||
# get_resume_checkpoint_until_find_ucp returns <resume_dir>/checkpoint-<step> with latest_universal,
|
||||
# or None; when None, no universal checkpoint exists and training starts from scratch.
|
||||
resume_checkpoint = trainer.get_resume_checkpoint_until_find_ucp()
|
||||
return resume_checkpoint
|
||||
|
||||
def train(self, trainer):
|
||||
logging_path = os.path.join(trainer.args.output_dir, 'logging.jsonl')
|
||||
logger.info(f'The logging file will be saved in: {logging_path}')
|
||||
resume_checkpoint = self._get_resume_checkpoint(trainer)
|
||||
try:
|
||||
trainer.train(resume_checkpoint)
|
||||
finally:
|
||||
res = self._save_trainer_state(trainer)
|
||||
if self.args.use_flash_ckpt and hasattr(trainer, 'flash_checkpointer'):
|
||||
trainer.wait_latest_checkpoint(trainer.FLASH_CKPT_WAIT_TIMEOUT, trainer.state.global_step)
|
||||
|
||||
return res
|
||||
|
||||
@staticmethod
|
||||
def _stat_dataset(dataset: Union[HfDataset, PackingDataset, LazyLLMDataset]):
|
||||
if isinstance(dataset, LazyLLMDataset):
|
||||
dataset = dataset.dataset
|
||||
if isinstance(dataset, HfDataset):
|
||||
lengths = dataset['lengths']
|
||||
lengths = [max(length) if isinstance(length, list) else length for length in lengths]
|
||||
else:
|
||||
lengths = dataset.packed_length
|
||||
_, stat_str = stat_array(lengths)
|
||||
logger.info(f'Dataset Token Length: {stat_str}')
|
||||
return stat_str
|
||||
|
||||
def _show_dataset(self, train_dataset, val_dataset):
|
||||
args = self.args
|
||||
predict_with_generate = getattr(args, 'predict_with_generate', False)
|
||||
if is_master():
|
||||
inputs = train_dataset[0] if hasattr(train_dataset, '__len__') else next(iter(train_dataset))
|
||||
if isinstance(inputs, list):
|
||||
inputs = inputs[0]
|
||||
self.template.print_inputs(inputs)
|
||||
elif hasattr(train_dataset, '__len__'):
|
||||
# Avoid the random mismatch issue in LazyLLMDataset.
|
||||
inputs = train_dataset[0]
|
||||
if val_dataset is not None and hasattr(val_dataset, '__len__') and len(val_dataset) == 0:
|
||||
val_dataset = None
|
||||
if not args.lazy_tokenize and not args.streaming:
|
||||
self.train_msg['train_dataset'] = self._stat_dataset(train_dataset)
|
||||
if val_dataset is not None and not predict_with_generate:
|
||||
self.train_msg['val_dataset'] = self._stat_dataset(val_dataset)
|
||||
|
||||
def _encode_dataset(self, train_dataset, val_dataset, pre_process=True):
|
||||
template = self.template
|
||||
args = self.args
|
||||
self._save_val_dataset(val_dataset)
|
||||
|
||||
predict_with_generate = getattr(args, 'predict_with_generate', False)
|
||||
datasets = [train_dataset, val_dataset]
|
||||
if not pre_process:
|
||||
return datasets
|
||||
|
||||
origin_template_model = template.model
|
||||
template.model = None # Avoid serializing the model.
|
||||
if args.truncation_strategy == 'split':
|
||||
if args.task_type != 'causal_lm' or template.mode != 'train' or args.use_chat_template:
|
||||
raise ValueError('`--truncation_strategy split` is currently only supported for pre-training.')
|
||||
assert not args.lazy_tokenize, '`--truncation_strategy split` does not support lazy_tokenize'
|
||||
|
||||
for i, dataset in enumerate(datasets):
|
||||
if dataset is None:
|
||||
continue
|
||||
if i == 1 and predict_with_generate:
|
||||
# val_dataset
|
||||
continue
|
||||
if not args.lazy_tokenize and not args.streaming:
|
||||
# Compatible with cached_dataset, only additionally write length here.
|
||||
preprocessor_cls = EncodePreprocessor if args.truncation_strategy == 'split' else AddLengthPreprocessor
|
||||
preprocessor = preprocessor_cls(template=template)
|
||||
batch_size = 100 if args.model_meta.is_multimodal else 1000
|
||||
dataset = preprocessor(
|
||||
dataset,
|
||||
num_proc=args.dataset_num_proc,
|
||||
load_from_cache_file=args.load_from_cache_file,
|
||||
strict=args.strict,
|
||||
batch_size=batch_size)
|
||||
if len(dataset) == 0:
|
||||
dataset = None
|
||||
datasets[i] = dataset
|
||||
template.model = origin_template_model
|
||||
|
||||
return datasets
|
||||
|
||||
|
||||
def sft_main(args: Optional[Union[List[str], SftArguments]] = None):
|
||||
return SwiftSft(args).main()
|
||||
@@ -0,0 +1,389 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import inspect
|
||||
import torch
|
||||
import transformers
|
||||
from packaging import version
|
||||
from peft.utils.other import ModulesToSaveWrapper
|
||||
from transformers import TrainingArguments
|
||||
from transformers.integrations import is_deepspeed_zero3_enabled
|
||||
from typing import List, Union
|
||||
|
||||
from swift.arguments import SftArguments
|
||||
from swift.trainers import calculate_max_steps
|
||||
from swift.tuner_plugin import Tuner, tuners_map
|
||||
from swift.tuners import Swift
|
||||
from swift.utils import (activate_parameters, find_all_linears, find_embedding, find_norm, freeze_parameters,
|
||||
get_logger, get_multimodal_target_regex)
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def apply_liger(model_type: str):
|
||||
try:
|
||||
from liger_kernel.transformers import (apply_liger_kernel_to_gemma, apply_liger_kernel_to_llama,
|
||||
apply_liger_kernel_to_mistral, apply_liger_kernel_to_mixtral,
|
||||
apply_liger_kernel_to_mllama, apply_liger_kernel_to_phi3,
|
||||
apply_liger_kernel_to_qwen2, apply_liger_kernel_to_qwen2_5_vl,
|
||||
apply_liger_kernel_to_qwen2_vl, apply_liger_kernel_to_qwen3)
|
||||
|
||||
from swift.model import ModelType
|
||||
if model_type in (ModelType.llama, ModelType.llama3, ModelType.llama3_1, ModelType.llama3_2):
|
||||
apply_liger_kernel_to_llama()
|
||||
elif model_type in (ModelType.mistral):
|
||||
apply_liger_kernel_to_mistral()
|
||||
elif model_type in (ModelType.mixtral):
|
||||
apply_liger_kernel_to_mixtral()
|
||||
elif model_type in (ModelType.gemma, ModelType.gemma2):
|
||||
apply_liger_kernel_to_gemma()
|
||||
elif model_type in (ModelType.gemma3_text):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_gemma3_text
|
||||
apply_liger_kernel_to_gemma3_text()
|
||||
elif model_type in (ModelType.gemma3_vision, ModelType.gemma3n):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_gemma3
|
||||
apply_liger_kernel_to_gemma3()
|
||||
elif model_type in (ModelType.qwen2, ModelType.qwen2_5):
|
||||
apply_liger_kernel_to_qwen2()
|
||||
elif model_type in (ModelType.qwen3, ModelType.qwen3_guard, ModelType.qwen3_thinking,
|
||||
ModelType.qwen3_nothinking, ModelType.qwen3_coder):
|
||||
apply_liger_kernel_to_qwen3()
|
||||
elif model_type in (ModelType.qwen3_moe, ModelType.qwen3_moe_thinking, ModelType.qwen3_coder):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_qwen3_moe
|
||||
apply_liger_kernel_to_qwen3_moe()
|
||||
elif model_type in (ModelType.qwen3_next, ModelType.qwen3_next_thinking):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_qwen3_next
|
||||
apply_liger_kernel_to_qwen3_next()
|
||||
elif model_type in (ModelType.phi3):
|
||||
apply_liger_kernel_to_phi3()
|
||||
elif model_type in (ModelType.llama3_2_vision):
|
||||
apply_liger_kernel_to_mllama()
|
||||
elif model_type in (ModelType.qwen2_vl):
|
||||
apply_liger_kernel_to_qwen2_vl()
|
||||
elif model_type in (ModelType.qwen2_5_vl, ModelType.qwen3_vl, ModelType.qwen3_vl_moe, ModelType.qvq):
|
||||
apply_liger_kernel_to_qwen2_5_vl()
|
||||
elif model_type in (ModelType.chatglm4, ModelType.glm4):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_glm4
|
||||
apply_liger_kernel_to_glm4()
|
||||
elif model_type in (ModelType.chatglm4v, ModelType.glm4v):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_glm4v
|
||||
apply_liger_kernel_to_glm4v()
|
||||
elif model_type in (ModelType.glm4v_moe):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_glm4v_moe
|
||||
apply_liger_kernel_to_glm4v_moe()
|
||||
elif model_type in (ModelType.internvl):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_internvl
|
||||
apply_liger_kernel_to_internvl()
|
||||
elif model_type in (ModelType.llama4):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_llama4
|
||||
apply_liger_kernel_to_llama4()
|
||||
elif model_type in (ModelType.llava1_5_hf, ModelType.llava_llama3_hf, ModelType.pixtral):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_llava
|
||||
apply_liger_kernel_to_llava()
|
||||
elif model_type in (ModelType.paligemma):
|
||||
from liger_kernel.transformers import apply_liger_kernel_to_paligemma
|
||||
apply_liger_kernel_to_paligemma()
|
||||
else:
|
||||
raise ValueError(f'Unsupported liger model_type: {model_type}')
|
||||
except ImportError:
|
||||
raise ImportError('Please upgrade liger-kernel to apply liger kernel to this model '
|
||||
'by running `pip install -U liger-kernel`')
|
||||
|
||||
|
||||
def get_target_modules(args, model) -> Union[str, List[str]]:
|
||||
"""Replace all-linear to actual modules"""
|
||||
if isinstance(args.target_modules, str):
|
||||
return args.target_modules
|
||||
target_modules = args.target_modules.copy()
|
||||
if 'all-linear' in target_modules:
|
||||
if model.model_meta.is_multimodal:
|
||||
return get_multimodal_target_regex(
|
||||
model,
|
||||
freeze_llm=args.freeze_llm,
|
||||
freeze_vit=args.freeze_vit,
|
||||
freeze_aligner=args.freeze_aligner,
|
||||
include_embedding='all-embedding' in target_modules)
|
||||
else:
|
||||
target_modules.remove('all-linear')
|
||||
target_modules += find_all_linears(model)
|
||||
if 'all-embedding' in target_modules:
|
||||
target_modules.remove('all-embedding')
|
||||
target_modules += find_embedding(model)
|
||||
return target_modules
|
||||
|
||||
|
||||
def get_modules_to_save(args, model, task_type=None):
|
||||
modules_to_save = args.modules_to_save.copy()
|
||||
if 'all-embedding' in args.modules_to_save:
|
||||
modules_to_save.remove('all-embedding')
|
||||
modules_to_save += find_embedding(model)
|
||||
if 'all-norm' in args.modules_to_save:
|
||||
modules_to_save.remove('all-norm')
|
||||
modules_to_save += find_norm(model)
|
||||
if task_type and task_type.lower() == 'seq_cls': # reward_model
|
||||
modules_to_save.append('v_head')
|
||||
return modules_to_save
|
||||
|
||||
|
||||
def get_vera_target_modules(model, config):
|
||||
"""This function is only useful on the vera tuner"""
|
||||
target_modules = config.target_modules
|
||||
modules_dict = {
|
||||
name: module.weight.shape
|
||||
for name, module in model.named_modules()
|
||||
if isinstance(module, torch.nn.Linear) and any([t in name for t in target_modules])
|
||||
} # only Linear for now
|
||||
if len(set(modules_dict.values())) > 1:
|
||||
v = [t for t in target_modules if 'v' in t]
|
||||
if not v:
|
||||
raise ValueError('Please manually pass in `vera_target_modules`, do not use `all-linear`,'
|
||||
'because Vera need all target linears to be the same size.')
|
||||
v = v[0]
|
||||
shape = [shape for name, shape in modules_dict.items() if v in name][0]
|
||||
names = [_name for _name, _shape in modules_dict.items() if _shape == shape]
|
||||
config.target_modules = [t for t in target_modules if any([t in name for name in names])]
|
||||
return config
|
||||
|
||||
|
||||
def prepare_adapter(args: SftArguments, model, *, template=None, train_dataset=None, task_type=None):
|
||||
from swift.tuners import (AdaLoraConfig, AdapterConfig, BOFTConfig, LLaMAProConfig, LongLoRAModelType, LoraConfig,
|
||||
LoRAConfig, ReftConfig, Swift, VeraConfig)
|
||||
task_type = (task_type or args.task_type).upper()
|
||||
target_modules = get_target_modules(args, model)
|
||||
modules_to_save = get_modules_to_save(args, model, task_type)
|
||||
lora_kwargs = {
|
||||
'r': args.lora_rank,
|
||||
'target_modules': target_modules,
|
||||
'lora_alpha': args.lora_alpha,
|
||||
'lora_dropout': args.lora_dropout,
|
||||
'bias': args.lora_bias,
|
||||
'modules_to_save': modules_to_save,
|
||||
'use_rslora': args.use_rslora,
|
||||
'use_dora': args.use_dora,
|
||||
'lorap_lr_ratio': args.lorap_lr_ratio,
|
||||
'init_lora_weights': args.init_weights,
|
||||
}
|
||||
if args.tuner_type in ('lora', 'longlora'):
|
||||
if args.use_swift_lora:
|
||||
lora_config = LoRAConfig(lora_dtype=args.lora_dtype, **lora_kwargs)
|
||||
model = Swift.prepare_model(model, lora_config)
|
||||
logger.info(f'lora_config: {lora_config}')
|
||||
elif args.tuner_backend == 'peft':
|
||||
if task_type == 'EMBEDDING':
|
||||
task_type = None
|
||||
elif task_type == 'RERANKER':
|
||||
task_type = 'SEQ_CLS'
|
||||
elif task_type == 'GENERATIVE_RERANKER':
|
||||
task_type = 'CAUSAL_LM'
|
||||
if args.target_parameters is not None:
|
||||
lora_kwargs['target_parameters'] = args.target_parameters
|
||||
lora_config = LoraConfig(task_type=task_type, lora_dtype=args.lora_dtype, **lora_kwargs)
|
||||
if args.init_weights == 'lora-ga':
|
||||
try:
|
||||
import lora_ga
|
||||
except ImportError as e:
|
||||
error_message = """
|
||||
Since 'LoRA-GA' is not implemented by PEFT, you will need to install it directly from GitHub.
|
||||
Command: 'pip install git+https://github.com/lxline/LoRA-GA.git'.
|
||||
"""
|
||||
logger.info(error_message)
|
||||
raise RuntimeError(error_message) from e
|
||||
model = lora_ga.entrypoint.get_lora_ga_model(
|
||||
model=model,
|
||||
data_collator=template.data_collator,
|
||||
dataset=train_dataset,
|
||||
batch_size=args.lora_ga_batch_size,
|
||||
num_iters=args.lora_ga_iters,
|
||||
max_length=args.lora_ga_max_length,
|
||||
direction=args.lora_ga_direction,
|
||||
dtype=args.lora_dtype,
|
||||
scale=args.lora_ga_scale,
|
||||
stable_gamma=args.lora_ga_stable_gamma,
|
||||
)
|
||||
else:
|
||||
model = Swift.prepare_model(model, lora_config)
|
||||
logger.info(f'lora_config: {lora_config}')
|
||||
elif args.tuner_backend == 'unsloth':
|
||||
if args.resume_from_checkpoint is None:
|
||||
if args.model_meta.is_multimodal:
|
||||
from unsloth import FastVisionModel as UnslothModel
|
||||
else:
|
||||
from unsloth import FastLanguageModel as UnslothModel
|
||||
assert args.tuner_type == 'lora', 'Unsloth does not support LongLoRA'
|
||||
lora_kwargs.pop('lorap_lr_ratio')
|
||||
model = UnslothModel.get_peft_model(
|
||||
model,
|
||||
use_gradient_checkpointing='unsloth',
|
||||
max_seq_length=args.max_length or 2048, # 2048 is the default value of unsloth
|
||||
**lora_kwargs,
|
||||
)
|
||||
logger.info(f'unsloth_config: {lora_kwargs}')
|
||||
if args.tuner_type == 'longlora':
|
||||
assert LongLoRAModelType.LLAMA in args.model_type
|
||||
assert version.parse(transformers.__version__) >= version.parse('4.39.3')
|
||||
from swift.tuners.longlora.llama import replace_llama_attn
|
||||
replace_llama_attn(model)
|
||||
model.config.group_size_ratio = 0.25
|
||||
elif args.tuner_type == 'adalora':
|
||||
lora_kwargs.pop('lorap_lr_ratio', None)
|
||||
lora_kwargs['rank_pattern'] = None
|
||||
adalora_config = AdaLoraConfig(
|
||||
task_type=task_type,
|
||||
**lora_kwargs,
|
||||
target_r=args.adalora_target_r,
|
||||
init_r=args.adalora_init_r,
|
||||
tinit=args.adalora_tinit,
|
||||
tfinal=args.adalora_tfinal,
|
||||
deltaT=args.adalora_deltaT,
|
||||
beta1=args.adalora_beta1,
|
||||
beta2=args.adalora_beta2,
|
||||
orth_reg_weight=args.adalora_orth_reg_weight,
|
||||
total_step=calculate_max_steps(args.training_args, train_dataset),
|
||||
)
|
||||
model = Swift.prepare_model(model, adalora_config)
|
||||
logger.info(f'adalora_config: {adalora_config}')
|
||||
elif args.tuner_type == 'llamapro':
|
||||
llamapro_config = LLaMAProConfig(
|
||||
model_type=model.model_meta.model_arch.arch_name,
|
||||
num_new_blocks=args.llamapro_num_new_blocks,
|
||||
num_groups=args.llamapro_num_groups)
|
||||
model = Swift.prepare_model(model, llamapro_config)
|
||||
logger.info(f'llamapro_config: {llamapro_config}')
|
||||
elif args.tuner_type == 'adapter':
|
||||
model_arch = model.model_meta.model_arch
|
||||
mlp_key = model_arch.mlp
|
||||
mlp_key = mlp_key.split('.{}.')[1]
|
||||
adapter_config = AdapterConfig(
|
||||
dim=model.config.hidden_size,
|
||||
target_modules=[mlp_key],
|
||||
hidden_pos=0,
|
||||
adapter_length=args.adapter_length,
|
||||
act_layer=args.adapter_act)
|
||||
model = Swift.prepare_model(model, adapter_config)
|
||||
logger.info(f'adapter_config: {adapter_config}')
|
||||
elif args.tuner_type == 'vera':
|
||||
vera_config = VeraConfig(
|
||||
r=args.vera_rank,
|
||||
target_modules=target_modules,
|
||||
projection_prng_key=args.vera_projection_prng_key,
|
||||
vera_dropout=args.vera_dropout,
|
||||
d_initial=args.vera_d_initial,
|
||||
modules_to_save=args.modules_to_save,
|
||||
)
|
||||
vera_config = get_vera_target_modules(model, vera_config)
|
||||
model = Swift.prepare_model(model, vera_config)
|
||||
logger.info(f'vera_config: {vera_config}')
|
||||
elif args.tuner_type == 'boft':
|
||||
boft_config = BOFTConfig(
|
||||
boft_block_size=args.boft_block_size,
|
||||
boft_block_num=args.boft_block_num,
|
||||
boft_n_butterfly_factor=args.boft_n_butterfly_factor,
|
||||
target_modules=target_modules,
|
||||
boft_dropout=args.boft_dropout,
|
||||
modules_to_save=args.modules_to_save,
|
||||
)
|
||||
model = Swift.prepare_model(model, boft_config)
|
||||
logger.info(f'boft_config: {boft_config}')
|
||||
elif args.tuner_type == 'fourierft':
|
||||
from peft import FourierFTConfig
|
||||
fourier_config = FourierFTConfig(
|
||||
target_modules=target_modules,
|
||||
modules_to_save=args.modules_to_save,
|
||||
n_frequency=args.fourier_n_frequency,
|
||||
scaling=args.fourier_scaling,
|
||||
)
|
||||
model = Swift.prepare_model(model, fourier_config)
|
||||
logger.info(f'fourier_config: {fourier_config}')
|
||||
elif args.tuner_type == 'reft':
|
||||
reft_config = ReftConfig(
|
||||
model_type=model.model_meta.model_arch,
|
||||
layer_key=args.reft_layer_key,
|
||||
r=args.reft_rank,
|
||||
layers=args.reft_layers,
|
||||
intervention_type=args.reft_intervention_type,
|
||||
args=args.reft_args,
|
||||
)
|
||||
logger.info(f'reft config: {reft_config}')
|
||||
model = Swift.prepare_model(model, {'reft': reft_config})
|
||||
elif args.tuner_type == 'bone':
|
||||
# Version loosing
|
||||
from peft import BoneConfig
|
||||
bone_config = BoneConfig(
|
||||
target_modules=target_modules,
|
||||
r=args.reft_rank,
|
||||
init_weights=args.init_weights,
|
||||
)
|
||||
logger.info(f'bone config: {bone_config}')
|
||||
model = Swift.prepare_model(model, bone_config)
|
||||
else:
|
||||
raise ValueError(f'Unknown tuner_type: {args.tuner_type}')
|
||||
return model
|
||||
|
||||
|
||||
def _patch_modules_to_save_zero3():
|
||||
if getattr(ModulesToSaveWrapper, '_patched', False):
|
||||
return
|
||||
ModulesToSaveWrapper._patched = True
|
||||
_old_setattr = ModulesToSaveWrapper.__setattr__
|
||||
|
||||
def _patched_setattr(self, name, value):
|
||||
_old_setattr(self, name, value)
|
||||
if name == 'ds_grads_remaining':
|
||||
for module in self.modules_to_save.values():
|
||||
module.ds_grads_remaining = value
|
||||
|
||||
ModulesToSaveWrapper.__setattr__ = _patched_setattr
|
||||
|
||||
|
||||
class TunerMixin:
|
||||
|
||||
@classmethod
|
||||
def prepare_model(cls, args, model, *, template=None, train_dataset=None, task_type=None):
|
||||
# transformers >= 4.45.0, apply liger in transformers https://github.com/huggingface/transformers/pull/32860
|
||||
# transformers < 4.45.0, apply liger in here
|
||||
if args.use_liger_kernel and 'use_liger_kernel' not in inspect.signature(TrainingArguments).parameters:
|
||||
# Apply liger
|
||||
apply_liger(args.model_type)
|
||||
|
||||
if args.is_adapter:
|
||||
if args.tuner_backend != 'unsloth' and args.tuner_type not in tuners_map:
|
||||
# Fix the name of the layer in xcomposer that contains Plora.
|
||||
# Unsloth prepares and loads lora outside this function when
|
||||
# resume_from_checkpoint, so do not disable grad here
|
||||
model.requires_grad_(False)
|
||||
if args.resume_from_checkpoint or args.adapters:
|
||||
if args.tuner_type in tuners_map:
|
||||
tuner: Tuner = tuners_map[args.tuner_type]
|
||||
else:
|
||||
tuner = Swift
|
||||
assert not args.adapters or len(args.adapters) == 1, f'args.adapters: {args.adapters}'
|
||||
model = tuner.from_pretrained(model, args.resume_from_checkpoint or args.adapters[0], is_trainable=True)
|
||||
else:
|
||||
if args.tuner_type in tuners_map:
|
||||
tuner: Tuner = tuners_map[args.tuner_type]
|
||||
model = tuner.prepare_model(args, model)
|
||||
else:
|
||||
model = prepare_adapter(
|
||||
args, model, template=template, train_dataset=train_dataset, task_type=task_type)
|
||||
# fix bug: Attempting to unscale FP16 gradients.
|
||||
# peft: https://github.com/huggingface/peft/issues/1249
|
||||
for p in model.parameters():
|
||||
if p.requires_grad and p.dtype == torch.float16:
|
||||
logger.info_once('Convert trainable parameters from fp16 to fp32.')
|
||||
p.data = p.data.to(dtype=torch.float32)
|
||||
elif args.tuner_type == 'full':
|
||||
model.train()
|
||||
model.requires_grad_(True)
|
||||
|
||||
freeze_parameters(model, args.freeze_parameters_ratio, args.freeze_parameters, args.freeze_parameters_regex)
|
||||
if args.trainable_parameters or args.trainable_parameters_regex:
|
||||
activate_parameters(model, args.trainable_parameters, args.trainable_parameters_regex)
|
||||
else:
|
||||
raise ValueError(f'args.tuner_type: {args.tuner_type}')
|
||||
|
||||
if args.use_galore:
|
||||
if args.galore_target_modules is None:
|
||||
args.galore_target_modules = find_all_linears(model)
|
||||
if args.galore_with_embedding:
|
||||
args.galore_target_modules += find_embedding(model)
|
||||
if is_deepspeed_zero3_enabled():
|
||||
_patch_modules_to_save_zero3()
|
||||
return model
|
||||
@@ -0,0 +1,82 @@
|
||||
# Copyright (c) ModelScope Contributors. All rights reserved.
|
||||
import numpy as np
|
||||
import os
|
||||
from datasets import load_from_disk
|
||||
|
||||
from swift.dataset import DatasetSyntax, sample_dataset
|
||||
from swift.template import update_generation_config_eos_token
|
||||
from swift.tuner_plugin import tuners_map
|
||||
from swift.tuners import Swift
|
||||
from swift.utils import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def prepare_adapter(args, model, adapters=None):
|
||||
if args.tuner_backend == 'unsloth':
|
||||
if args.model_meta.is_multimodal:
|
||||
from unsloth import FastVisionModel as UnslothModel
|
||||
else:
|
||||
from unsloth import FastLanguageModel as UnslothModel
|
||||
UnslothModel.for_inference(model)
|
||||
return model
|
||||
if args.tuner_type in tuners_map:
|
||||
tuner = tuners_map[args.tuner_type]
|
||||
else:
|
||||
tuner = Swift
|
||||
# compat deploy
|
||||
adapters = adapters if adapters is not None else args.adapters
|
||||
for adapter in adapters:
|
||||
model = tuner.from_pretrained(model, adapter)
|
||||
if args.tuner_type == 'bone':
|
||||
# Bone has a problem of float32 matmul with bloat16 in `peft==0.14.0`
|
||||
model.to(model.dtype)
|
||||
return model
|
||||
|
||||
|
||||
def prepare_model_template(args, **kwargs):
|
||||
adapters = kwargs.get('adapters')
|
||||
model, processor = args.get_model_processor(**kwargs)
|
||||
template = args.get_template(processor)
|
||||
if model is not None:
|
||||
if template.use_model:
|
||||
template.model = model
|
||||
model = prepare_adapter(args, model, adapters=adapters)
|
||||
if args.task_type == 'causal_lm':
|
||||
update_generation_config_eos_token(model.generation_config, template)
|
||||
return model, template
|
||||
|
||||
|
||||
def _select_dataset(args, dataset):
|
||||
if 'length' in dataset.column_names and 'lengths' not in dataset.column_names:
|
||||
# Compatible with ms-swift 3.x cache_dataset
|
||||
dataset = dataset.rename_column('length', 'lengths')
|
||||
max_length = args.max_length
|
||||
if args.truncation_strategy == 'delete':
|
||||
lengths = dataset['lengths']
|
||||
idxs = [
|
||||
i for i, length in enumerate(lengths) if (max(length) if isinstance(length, list) else length) <= max_length
|
||||
]
|
||||
new_dataset = dataset.select(idxs)
|
||||
else:
|
||||
new_dataset = dataset
|
||||
if len(new_dataset) < len(dataset):
|
||||
logger.info(f'Dataset filtered, origin length: {len(dataset)}, filtered dataset length: {len(new_dataset)}')
|
||||
return new_dataset
|
||||
|
||||
|
||||
def get_cached_dataset(args):
|
||||
train_datasets, val_datasets = [], []
|
||||
random_state = np.random.RandomState(args.data_seed)
|
||||
for cached_dataset, datasets in zip([args.cached_dataset, args.cached_val_dataset], [train_datasets, val_datasets]):
|
||||
for path in cached_dataset:
|
||||
if os.path.exists(path):
|
||||
dataset_sample = None
|
||||
else:
|
||||
path, dataset_sample = DatasetSyntax._safe_split(path, '#', True, 'right')
|
||||
dataset = _select_dataset(args, load_from_disk(path))
|
||||
if dataset_sample is not None:
|
||||
dataset = sample_dataset(
|
||||
dataset, int(dataset_sample), args.dataset_shuffle, random_state=random_state, shuffle_all=True)
|
||||
datasets.append(dataset)
|
||||
return train_datasets, val_datasets
|
||||
Reference in New Issue
Block a user