diff --git a/.github/workflows/CI.yml b/.github/workflows/CI.yml index a3d0c3ce..a654ac46 100644 --- a/.github/workflows/CI.yml +++ b/.github/workflows/CI.yml @@ -5,12 +5,14 @@ on: branches: - main - master + - learn2code tags: - "*" pull_request: branches: - main - master + - learn2code workflow_dispatch: permissions: @@ -56,7 +58,7 @@ jobs: run: | set +e has_updated=$(git diff --name-only '${{ github.event.pull_request.base.sha }}' | grep -E 'pyproject.toml|Cargo.toml') - is_main_or_release='${{ github.ref == 'refs/heads/main' || github.ref == 'refs/heads/master' || startsWith(github.ref, 'refs/tags/') }}' + is_main_or_release='${{ github.ref == 'refs/heads/main' || github.ref == 'refs/heads/master' || github.ref == 'refs/heads/learn2code' || startsWith(github.ref, 'refs/tags/') }}' if [[ -n "${has_updated}" ]] || [[ "${is_main_or_release}" == 'true' ]]; then echo "should_build=true" >> $GITHUB_OUTPUT else @@ -202,7 +204,7 @@ jobs: sudo apt-get update sudo apt-get install --yes --upgrade build-essential cmake protobuf-compiler libssl-dev glibc-source musl-tools - name: Upload wheels - uses: actions/upload-artifact@v4.4 + uses: actions/upload-artifact@v4.4.1 with: overwrite: true name: release-wheel-linux-${{ matrix.target }}-${{ github.run_id }} @@ -228,7 +230,7 @@ jobs: target: ${{ matrix.target }} args: --release --out dist --find-interpreter - name: Upload wheels - uses: actions/upload-artifact@v4.4 + uses: actions/upload-artifact@v4.4.1 with: overwrite: true name: release-wheel-windows-${{ matrix.target }}-${{ github.run_id }} @@ -253,7 +255,7 @@ jobs: target: ${{ matrix.target }} args: --release --out dist --find-interpreter - name: Upload wheels - uses: actions/upload-artifact@v4.4 + uses: actions/upload-artifact@v4.4.1 with: overwrite: true name: release-wheel-macos-${{ matrix.target }}-${{ github.run_id }} @@ -271,7 +273,7 @@ jobs: command: sdist args: --out dist - name: Upload sdist - uses: actions/upload-artifact@v4.4 + uses: actions/upload-artifact@v4.4.1 with: overwrite: true name: release-sdist-${{ github.run_id }} @@ -281,7 +283,7 @@ jobs: name: Release runs-on: ubuntu-latest if: "startsWith(github.ref, 'refs/tags/')" - needs: [build-linux, build-windows, build-macos, sdist] + needs: [build-linux, build-windows, build-macos, sdist, tests] steps: - uses: actions/download-artifact@v4 with: diff --git a/Cargo.toml b/Cargo.toml index 97bb1367..7537edb5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "dolma" -version = "1.1.1" +version = "1.2.0-dev7" edition = "2021" license = "Apache-2.0" diff --git a/pyproject.toml b/pyproject.toml index 86b03239..45ebf886 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "dolma" -version = "1.1.2" +version = "1.2.0.dev7" description = "Data filters" license = { text = "Apache-2.0" } readme = "README.md" diff --git a/python/dolma/taggers/code/__init__.py b/python/dolma/taggers/code/__init__.py index 13117d6c..cc23eea2 100644 --- a/python/dolma/taggers/code/__init__.py +++ b/python/dolma/taggers/code/__init__.py @@ -4,6 +4,7 @@ CodeSecretsTagger, CodeStarCoderTaggers, CodeStarCoderTaggers2, + Learn2CodeTaggers, ) __all__ = [ @@ -12,4 +13,5 @@ "CodeRedPajamaTaggers", "CodeStarCoderTaggers", "CodeStarCoderTaggers2", + "Learn2CodeTaggers", ] diff --git a/python/dolma/taggers/code/code_taggers.py b/python/dolma/taggers/code/code_taggers.py index 31b57087..2cac0222 100644 --- a/python/dolma/taggers/code/code_taggers.py +++ b/python/dolma/taggers/code/code_taggers.py @@ -21,10 +21,16 @@ if CODE_DEPENDENCIES_AVAILABLE: from .starcoder import get_nl_ratio from .utils import ( + b64_filter, filter_html, get_ext_to_lang_mapping, + get_line_stats, + get_proportion_alphabetic_chars, get_secrets, get_whitespace_regex, + hexadecimal_filter, + special_text_file_filter, + unicode_filter, ) @@ -269,3 +275,108 @@ def predict(self, doc: DocumentWithMetadata) -> DocResult: # type: ignore spans.append(Span(start=0, end=doc_length, type="code_to_text_ratio_html_doc", score=code_to_text_ratio)) return DocResult(doc=doc, spans=spans) + + +@TaggerRegistry.add("learn2code_taggers_v1") +class Learn2CodeTaggers(BaseTaggerWithMetadata): + """ + Based on a mix of filters from StarCoder and Granite + """ + + def __init__(self) -> None: + check_code_dependencies() + self.ext_to_lang_mapping = get_ext_to_lang_mapping() + super().__init__() + + def predict(self, doc: DocumentWithMetadata) -> DocResult: # type: ignore + spans: List[Span] = [] + doc_length = len(doc.text) + + num_github_stars = doc.metadata.get("max_stars_count", 0) or doc.metadata.get("star_events_count", 0) or 0 + proportion_alpha = get_proportion_alphabetic_chars(doc.text) + has_xml_template = 1.0 if " StarCoderRegexFilterResults: + all_matches = re.findall(regex_string, text) + + match_lengths = [len(match) for match in all_matches] + longest_match = max(match_lengths) if match_lengths else 0 + proportion_match = sum(match_lengths) / len(text) + + return StarCoderRegexFilterResults(longest_match=longest_match, proportion_match=proportion_match) + + +def b64_filter(text: str) -> StarCoderRegexFilterResults: + """ + Taken from the StarCoder2 paper. + """ + regex = r"[a-zA-Z0-9+/\n=]{64,}" + return regex_match(regex, text) + + +def hexadecimal_filter(text: str) -> StarCoderRegexFilterResults: + """ + Taken from StarCoder2 paper. + The escaped literal case, e.g. "\\x48\\x31\\xc0\\x50\\x68\\x2f\\x2f\\x73\\x68", + is a bit broken, because it'll always drop the first byte in the sequence due to + how \b is interpreted in that context. + """ + regex = r"(?:\b(?:0x|\\x)?[0-9a-fA-F]{2}(?:,|\b\s*)){8,}" + return regex_match(regex, text) + + +def unicode_filter(text: str) -> StarCoderRegexFilterResults: + """ + Taken from the StarCoder2 paper. + """ + regex = r"(?:\\u[0-9a-fA-F]{4}){8,}" + return regex_match(regex, text) + + +def get_proportion_alphabetic_chars(text: str) -> float: + """Calculates the proportion of characters in passed text that are alphabetic""" + nonalpha = re.sub(r"[^A-Za-z]", "", text) + return len(nonalpha) / len(text) + + +@dataclass +class LineStats: + total_count: int + mean_length: float + max_length: int + + +def get_line_stats(text: str) -> LineStats: + """Finds some summary stats about the lines in the passed text""" + + lines = text.split("\n") + line_lengths = [len(line) for line in lines] + + return LineStats( + total_count=len(lines), mean_length=sum(line_lengths) / len(lines), max_length=max(line_lengths) + ) + + def filter_html(html: str) -> float: """Filter HTML files based on displayed text VS code ratio""" try: @@ -80,3 +149,16 @@ def get_ext_to_lang_mapping() -> Dict[str, str]: path = Path(__file__).parent / "../../data/ext_to_lang_mapping.json" with smart_open.open(path, "r") as f: return json.load(f) + + +def special_text_file_filter(filepath: str, lang: str) -> bool: + if lang == "text": # TODO: include markdown as well? + filename = Path(filepath).stem.lower() + + if "requirement" in filename: + return True + + if filename in {"readme", "todo", "description", "cmakelists"}: + return True + + return False diff --git a/tests/python/test_code.py b/tests/python/test_code.py index 6af2738f..13cd1b05 100644 --- a/tests/python/test_code.py +++ b/tests/python/test_code.py @@ -6,8 +6,11 @@ """ +import base64 import re import unittest +from pathlib import Path +from typing import Optional from bs4 import BeautifulSoup @@ -17,6 +20,7 @@ CodeRedPajamaTaggers, CodeSecretsTagger, CodeStarCoderTaggers2, + Learn2CodeTaggers, ) DOC_WITH_SECRETS_AND_COPYRIGHT = """ @@ -236,3 +240,184 @@ def test_html_tagger_doc_too_short(self): self.assertEqual(result.spans[3].type, "code_to_text_ratio_html_doc") self.assertEqual(result.spans[3].score, 0.0) + + +class TestLearn2CodeTaggers(unittest.TestCase): + def mk_doc(self, id="0", text="asdf", source=__file__, metadata_overrides=None) -> DocumentWithMetadata: + return DocumentWithMetadata( + id=id, + text=text, + source=source, + metadata={ + "ext": "py", + "path": "/some/path/to/this/file.py", + **(metadata_overrides or {}), + }, + ) + + def assert_doc_span_score( + self, doc: DocumentWithMetadata, span_type: str, span_score: float, msg: Optional[str] = None + ) -> None: + for span in doc.spans: + if span.type is span_type: + self.assertEqual(span.score, span_score, msg=msg) + break + else: + raise AssertionError(f"No span of type `{span_type}`") + + def test_doc_length(self) -> None: + tagger = Learn2CodeTaggers() + doc = self.mk_doc(text="a" * 9001) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score(tagged_doc, "num_chars_doc", 9001) + + def test_gh_star_tagging(self) -> None: + tagger = Learn2CodeTaggers() + + doc1 = self.mk_doc(metadata_overrides={"max_stars_count": 9000}) + tagged_doc1 = tagger.predict(doc1) + self.assert_doc_span_score(tagged_doc1, "num_github_stars_doc", 9000) + + doc2 = self.mk_doc(metadata_overrides={"star_events_count": 9000}) + tagged_doc2 = tagger.predict(doc2) + self.assert_doc_span_score(tagged_doc2, "num_github_stars_doc", 9000) + + def test_proportion_alpha_doc(self) -> None: + tagger = Learn2CodeTaggers() + + texts_percentage_pairs = [ + ("12345!@#%$", 0.0), + ("f234567890", 0.1), + ("f2a4567890", 0.2), + ("fDSa5j7890", 0.5), + ("ABcdEFgHij", 1.0), + ] + + for doc_text, expected_percent_alpha in texts_percentage_pairs: + doc = self.mk_doc(text=doc_text) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score(tagged_doc, "proportion_alpha_doc", expected_percent_alpha) + + def test_has_xml_template_doc(self) -> None: + tagger = Learn2CodeTaggers() + + not_xml = "a" * 100 + "\n' + is_xml_doc = self.mk_doc(text=is_xml) + tagged_is_xml_doc = tagger.predict(is_xml_doc) + self.assert_doc_span_score(tagged_is_xml_doc, "has_xml_template_doc", 1.0) + + def test_line_stats(self) -> None: + tagger = Learn2CodeTaggers() + + lines = ["1" * 50, "2" * 150, "3" * 100] + text = "\n".join(lines) + + doc = self.mk_doc(text=text) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score(tagged_doc, "num_lines_doc", 3) + self.assert_doc_span_score(tagged_doc, "mean_line_length_doc", 100.0) + self.assert_doc_span_score(tagged_doc, "max_line_length_doc", 150) + + def test_b64_encoding_tagging(self) -> None: + tagger = Learn2CodeTaggers() + + to_encode = b"hello i am some text that is sufficiently long to encode and trigger the regexes we have" + base64_encoded_substring = base64.b64encode(to_encode).decode("utf-8") + full_text = ( + " " * len(base64_encoded_substring) + + base64_encoded_substring + + " " * 2 * len(base64_encoded_substring) + ) + + doc = self.mk_doc(text=full_text) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score(tagged_doc, "longest_seq_b64_doc", len(base64_encoded_substring)) + self.assert_doc_span_score(tagged_doc, "proportion_b64_doc", 0.25) + + def test_hexadecimal_filter(self) -> None: + tagger = Learn2CodeTaggers() + + hex_examples = [ + ("c_style_array", "0x1A,0x2B,0x3C,0x4D,0x5E,0x6F,0x7A,0x8B "), + ("hex_dump", "1A 2B 3C 4D 5E 6F 7A 8B "), + ] + + for name, example in hex_examples: + full_text = "a" * (len(example) - 1) + " " + example + "z" * 2 * len(example) + doc = self.mk_doc(text=full_text) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score( + tagged_doc, "longest_seq_hexadecimal_doc", len(example), f"{name} is wrong length" + ) + self.assert_doc_span_score(tagged_doc, "proportion_hexadecimal_doc", 0.25, f"{name} wrong percentage") + + def test_unicode_filter(self) -> None: + tagger = Learn2CodeTaggers() + unicode = "\\u0041\\u0042\\u0043\\u0044\\u0045\\u0046\\u0047\\u0048" + full_text = " " * len(unicode) + unicode + " " * len(unicode) + unicode + doc = self.mk_doc(text=full_text) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score(tagged_doc, "longest_seq_unicode_doc", len(unicode)) + self.assert_doc_span_score(tagged_doc, "proportion_unicode_doc", 0.5) + + def test_proportion_comments_doc(self) -> None: + tagger = Learn2CodeTaggers() + text = """ + def thingy(asdf): + '''I am a docstring''' + return asdf + """.strip() + + expected_comment_length = len("I am a docstring") + + doc = self.mk_doc(text=text, metadata_overrides={"ext": "py"}) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score(tagged_doc, "proportion_comments_doc", expected_comment_length / len(text)) + + def test_proportion_text_in_html(self) -> None: + tagger = Learn2CodeTaggers() + visible_text = "I am the visible text" * 20 + text = f"""
{visible_text}
""".strip() + + expected_visible_text_length = len(visible_text) + doc = self.mk_doc(text=text, metadata_overrides={"ext": "html"}) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score( + tagged_doc, "proportion_text_in_html_doc", expected_visible_text_length / len(text) + ) + + def test_is_special_text_file_doc(self) -> None: + tagger = Learn2CodeTaggers() + + special_ones = [ + "asdf/requirement.txt", + "asdf/requirements.txt", + "asdf/dev_requirements.txt", + "asdf/readme.txt", + "asdf/todo.txt", + "asdf/Description.txt", + "asdf/CMAKELISTS.txt", + ] + + not_so_special_ones = ["asdf/birthdays.txt", "asdf/readme.md", "readme"] + + for special_one in special_ones: + extension = Path(special_one).suffix.replace(".", "") + doc = self.mk_doc(special_one, metadata_overrides={"path": special_one, "ext": extension}) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score( + tagged_doc, "is_special_text_file_doc", 1, msg=f"`{special_one}` wasn't so special" + ) + + for not_special in not_so_special_ones: + extension = Path(not_special).suffix.replace(".", "") + doc = self.mk_doc(not_special, metadata_overrides={"path": not_special, "ext": extension}) + tagged_doc = tagger.predict(doc) + self.assert_doc_span_score( + tagged_doc, "is_special_text_file_doc", 0, msg=f"`{not_special}` was seen as special" + )