chore: import upstream snapshot with attribution
Lint test / lint (push) Has been cancelled

This commit is contained in:
wehub-resource-sync
2026-07-13 13:34:58 +08:00
commit a203934033
1368 changed files with 175001 additions and 0 deletions
+41
View File
@@ -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={},
)
+1
View File
@@ -0,0 +1 @@
from .app import SwiftApp, app_main
+43
View File
@@ -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()
+137
View File
@@ -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 = {'<': '&lt;', '>': '&gt;', '*': '&ast;'}
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
+23
View File
@@ -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': '📁 上传'
}
}
+58
View File
@@ -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
+2
View File
@@ -0,0 +1,2 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
from .eval import SwiftEval, eval_main
+160
View File
@@ -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()
+267
View File
@@ -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,
)
+6
View File
@@ -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
+47
View File
@@ -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()
+54
View File
@@ -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()
+61
View File
@@ -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 = []
+73
View File
@@ -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')
+292
View File
@@ -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()
+26
View File
@@ -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={},
)
+289
View File
@@ -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.')
+312
View File
@@ -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
+116
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
from .sampling import sampling_main
+59
View File
@@ -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
+160
View File
@@ -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
+104
View File
@@ -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()
+80
View File
@@ -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
+233
View File
@@ -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
+5
View File
@@ -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
+78
View File
@@ -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
+17
View File
@@ -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()
+248
View File
@@ -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()
+341
View File
@@ -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()
+389
View File
@@ -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
+82
View File
@@ -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