Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 92 additions & 7 deletions src/mcore_bridge/config/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,17 +16,102 @@

logger = get_logger()

# Layer-pattern strings are list literals with repetition, for example
# "([0]*3+[1]*1)*3". The alphabet is tiny, but eval still accepts "[0]*10**9"
# and will try to allocate it. Parse the grammar instead and refuse a result
# longer than a model would ever have.
_MAX_PATTERN_ITEMS = 100_000


# code borrowed from NVIDIA/Megatron-LM
def _eval_pattern(pattern):
""" Validate and evaluate a string containing a Python list expression """
"""Parse a layer-pattern list expression. No eval."""
assert isinstance(pattern, str)

# validate input, only allow comma, digits, [, ], (, ), +, and *
if bool(re.compile(r'[^,\d\[\]\(\)\+\*]').search(pattern)):
compact = re.sub(r'\s+', '', pattern)
if compact == '' or re.search(r'[^,\d\[\]\(\)\+\*]', compact):
raise ValueError(f'Invalid pattern: {pattern}')

return eval(pattern)
if '**' in compact:
raise ValueError(f'Invalid pattern: {pattern}')
parser = _PatternParser(compact)
value = parser.parse()
if parser.i != len(compact):
raise ValueError(f'Invalid pattern: {pattern}')
if not isinstance(value, list):
raise ValueError(f'Invalid pattern: {pattern}')
return value


class _PatternParser:

def __init__(self, text):
self.text = text
self.i = 0

def parse(self):
return self._expr()

def _peek(self):
return self.text[self.i] if self.i < len(self.text) else ''

def _eat(self, char):
if self._peek() != char:
raise ValueError(f'Invalid pattern: {self.text}')
self.i += 1

def _expr(self):
value = self._term()
while self._peek() == '+':
self._eat('+')
right = self._term()
if not isinstance(value, list) or not isinstance(right, list):
raise ValueError(f'Invalid pattern: {self.text}')
value = value + right
if len(value) > _MAX_PATTERN_ITEMS:
raise ValueError(f'Pattern is longer than {_MAX_PATTERN_ITEMS}: {self.text}')
return value

def _term(self):
value = self._atom()
if self._peek() == '*':
self._eat('*')
count = self._number()
if not isinstance(value, list):
raise ValueError(f'Invalid pattern: {self.text}')
if count > _MAX_PATTERN_ITEMS or len(value) * count > _MAX_PATTERN_ITEMS:
raise ValueError(f'Pattern is longer than {_MAX_PATTERN_ITEMS}: {self.text}')
value = value * count
return value

def _atom(self):
if self._peek() == '(':
self._eat('(')
value = self._expr()
self._eat(')')
return value
if self._peek() == '[':
return self._list()
raise ValueError(f'Invalid pattern: {self.text}')

def _list(self):
self._eat('[')
if self._peek() == ']':
self._eat(']')
return []
items = [self._number()]
while self._peek() == ',':
self._eat(',')
items.append(self._number())
if len(items) > _MAX_PATTERN_ITEMS:
raise ValueError(f'Pattern is longer than {_MAX_PATTERN_ITEMS}: {self.text}')
self._eat(']')
return items

def _number(self):
start = self.i
if not self._peek().isdigit():
raise ValueError(f'Invalid pattern: {self.text}')
while self._peek().isdigit():
self.i += 1
return int(self.text[start:self.i])


# code borrowed from NVIDIA/Megatron-LM
Expand Down
35 changes: 35 additions & 0 deletions tests/test_layer_pattern.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
"""Layer patterns are parsed, not eval'd, and huge repetitions are refused."""
import ast
from pathlib import Path


def _parser():
path = Path(__file__).resolve().parents[1] / 'src/mcore_bridge/config/model_config.py'
tree = ast.parse(path.read_text())
body = [
node for node in tree.body if (isinstance(node, ast.Assign) and any(
isinstance(target, ast.Name) and target.id == '_MAX_PATTERN_ITEMS' for target in node.targets)) or (
isinstance(node, (ast.FunctionDef, ast.ClassDef)) and node.name in {'_eval_pattern', '_PatternParser'})
]
module = ast.Module(body=body, type_ignores=[])
ast.fix_missing_locations(module)
namespace = {'re': __import__('re')}
exec(compile(module, str(path), 'exec'), namespace)
return namespace['_eval_pattern']


def test_documented_pattern_matches_the_expanded_list():
parse = _parser()
assert parse('([0]*3+[1]*1)*3') == [0, 0, 0, 1] * 3
assert parse('([1]+[0]*2)') == [1, 0, 0]


def test_exponent_and_huge_repeat_are_rejected():
parse = _parser()
for pattern in ('[0]*10**9', '[0]*(10**9)', '[0]*1000000'):
try:
parse(pattern)
except ValueError:
continue
raise AssertionError(pattern)