From 7064f796d549db0f03142e08e582eb387a99a775 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Fri, 25 Sep 2026 20:20:01 -0400 Subject: [PATCH 01/12] Performance updates --- CMakeLists.txt | 3 + codon/cir/llvm/llvisitor.cpp | 18 +- codon/cir/llvm/llvm.h | 1 + codon/cir/llvm/optimize.cpp | 199 +++++- codon/cir/transform/manager.cpp | 7 + codon/cir/transform/pythonic/enumerate.cpp | 178 ++++++ codon/cir/transform/pythonic/enumerate.h | 22 + codon/cir/transform/pythonic/generator.cpp | 231 +++++++ codon/cir/transform/pythonic/generator.h | 19 + stdlib/algorithms/timsort.codon | 8 + stdlib/heapq.codon | 38 +- stdlib/internal/__init_test__.codon | 8 +- stdlib/internal/builtin.codon | 121 ++-- stdlib/internal/python.codon | 4 +- stdlib/internal/types/collections/list.codon | 22 +- stdlib/internal/types/generator.codon | 26 +- stdlib/internal/types/str.codon | 12 + stdlib/itertools.codon | 95 ++- stdlib/math.codon | 33 +- stdlib/re.codon | 26 +- test/cir/llvm/optimize.cpp | 608 ++++++++++++++++++- test/cir/transform/enumerate.cpp | 189 ++++++ test/core/bltin.codon | 48 +- test/core/containers.codon | 37 ++ test/core/generators.codon | 376 ++++++++++++ test/core/numerics.codon | 16 + test/main.cpp | 1 + test/numpy/test_indexing.codon | 48 ++ test/parser/llvm.codon | 10 +- test/parser/typecheck/test_function.codon | 11 +- test/parser/typecheck/test_infer.codon | 4 +- test/python/pyext.py | 14 + test/stdlib/itertools_test.codon | 110 ++++ test/stdlib/llvm_test.codon | 10 +- test/stdlib/re_test.codon | 12 + test/stdlib/sort_test.codon | 29 + test/stdlib/str_test.codon | 11 + test/transform/enumerate.codon | 224 +++++++ 38 files changed, 2621 insertions(+), 208 deletions(-) create mode 100644 codon/cir/transform/pythonic/enumerate.cpp create mode 100644 codon/cir/transform/pythonic/enumerate.h create mode 100644 test/cir/transform/enumerate.cpp create mode 100644 test/transform/enumerate.codon diff --git a/CMakeLists.txt b/CMakeLists.txt index 77575e3ee..ce239496d 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -265,6 +265,7 @@ set(CODON_HPPFILES codon/cir/transform/parallel/schedule.h codon/cir/transform/pass.h codon/cir/transform/pythonic/dict.h + codon/cir/transform/pythonic/enumerate.h codon/cir/transform/pythonic/format.h codon/cir/transform/pythonic/generator.h codon/cir/transform/pythonic/io.h @@ -375,6 +376,7 @@ set(CODON_CPPFILES codon/cir/transform/parallel/schedule.cpp codon/cir/transform/pass.cpp codon/cir/transform/pythonic/dict.cpp + codon/cir/transform/pythonic/enumerate.cpp codon/cir/transform/pythonic/format.cpp codon/cir/transform/pythonic/generator.cpp codon/cir/transform/pythonic/io.cpp @@ -537,6 +539,7 @@ set(CODON_TEST_CPPFILES test/cir/instr.cpp test/cir/llvm/optimize.cpp test/cir/module.cpp + test/cir/transform/enumerate.cpp test/cir/transform/format.cpp test/cir/transform/manager.cpp test/cir/type.cpp diff --git a/codon/cir/llvm/llvisitor.cpp b/codon/cir/llvm/llvisitor.cpp index 1ff1fef26..bf2be42ff 100644 --- a/codon/cir/llvm/llvisitor.cpp +++ b/codon/cir/llvm/llvisitor.cpp @@ -2616,6 +2616,7 @@ void LLVMVisitor::visit(const ForFlow *x) { auto *condBlock = llvm::BasicBlock::Create(*context, "for.cond", func); auto *bodyBlock = llvm::BasicBlock::Create(*context, "for.body", func); + auto *checkBlock = llvm::BasicBlock::Create(*context, "for.check", func); auto *cleanupBlock = llvm::BasicBlock::Create(*context, "for.cleanup", func); auto *exitBlock = llvm::BasicBlock::Create(*context, "for.exit", func); @@ -2633,7 +2634,10 @@ void LLVMVisitor::visit(const ForFlow *x) { process(x->getIter()); auto *iter = value; B->SetInsertPoint(block); - B->CreateBr(condBlock); + B->CreateCondBr(B->CreateCall(coroDone, iter), exitBlock, condBlock); + + B->SetInsertPoint(checkBlock); + B->CreateCondBr(B->CreateCall(coroDone, iter), exitBlock, condBlock); block = condBlock; call(coroResume, {iter}); @@ -2653,11 +2657,11 @@ void LLVMVisitor::visit(const ForFlow *x) { block = bodyBlock; enterLoop( - {/*breakBlock=*/exitBlock, /*continueBlock=*/condBlock, /*loopId=*/x->getId()}); + {/*breakBlock=*/exitBlock, /*continueBlock=*/checkBlock, /*loopId=*/x->getId()}); process(x->getBody()); exitLoop(); B->SetInsertPoint(block); - B->CreateBr(condBlock); + B->CreateBr(checkBlock); B->SetInsertPoint(cleanupBlock); B->CreateCall(coroDestroy, iter); @@ -3139,6 +3143,7 @@ void LLVMVisitor::codegenPipeline( auto *condBlock = llvm::BasicBlock::Create(*context, "pipeline.cond", func); auto *bodyBlock = llvm::BasicBlock::Create(*context, "pipeline.body", func); + auto *checkBlock = llvm::BasicBlock::Create(*context, "pipeline.check", func); auto *cleanupBlock = llvm::BasicBlock::Create(*context, "pipeline.cleanup", func); auto *exitBlock = llvm::BasicBlock::Create(*context, "pipeline.exit", func); @@ -3155,7 +3160,10 @@ void LLVMVisitor::codegenPipeline( auto *iter = value; B->SetInsertPoint(block); - B->CreateBr(condBlock); + B->CreateCondBr(B->CreateCall(coroDone, iter), exitBlock, condBlock); + + B->SetInsertPoint(checkBlock); + B->CreateCondBr(B->CreateCall(coroDone, iter), exitBlock, condBlock); block = condBlock; call(coroResume, {iter}); @@ -3175,7 +3183,7 @@ void LLVMVisitor::codegenPipeline( codegenPipeline(stages, where + 1); B->SetInsertPoint(block); - B->CreateBr(condBlock); + B->CreateBr(checkBlock); B->SetInsertPoint(cleanupBlock); B->CreateCall(coroDestroy, iter); diff --git a/codon/cir/llvm/llvm.h b/codon/cir/llvm/llvm.h index b562aad6a..23ed92e3e 100644 --- a/codon/cir/llvm/llvm.h +++ b/codon/cir/llvm/llvm.h @@ -18,6 +18,7 @@ #include "llvm/Analysis/RegionPass.h" #include "llvm/Analysis/TargetLibraryInfo.h" #include "llvm/Analysis/TargetTransformInfo.h" +#include "llvm/Analysis/ValueTracking.h" #include "llvm/AsmParser/Parser.h" #include "llvm/Bitcode/BitcodeWriter.h" #include "llvm/CodeGen/CommandFlags.h" diff --git a/codon/cir/llvm/optimize.cpp b/codon/cir/llvm/optimize.cpp index a492f3b77..dd0d29b00 100644 --- a/codon/cir/llvm/optimize.cpp +++ b/codon/cir/llvm/optimize.cpp @@ -1109,6 +1109,78 @@ struct AllocationAutoFree : public llvm::PassInfoMixin { /// function pointer comparisons. This pass puts them into a somewhat /// easier-to-analyze form. struct CoroBranchSimplifier : public llvm::PassInfoMixin { + static bool mergeGuardedState(llvm::Loop &loop) { + llvm::SmallVector states; + for (auto &phi : loop.getHeader()->phis()) + if (phi.getType()->isPointerTy()) + states.push_back(&phi); + if (states.size() > 8) + return false; + for (unsigned first = 0; first < states.size(); ++first) { + for (unsigned second = first + 1; second < states.size(); ++second) { + auto *leftRoot = states[first]; + auto *rightRoot = states[second]; + llvm::SmallVector, 16> visited; + std::function + equivalent = [&](llvm::Value *left, llvm::Value *right, + llvm::BasicBlock *predecessor, + llvm::BasicBlock *successor) { + if (left == right) + return true; + if (predecessor) { + auto *branch = + llvm::dyn_cast(predecessor->getTerminator()); + auto *compare = + branch && branch->isConditional() + ? llvm::dyn_cast(branch->getCondition()) + : nullptr; + if (compare && compare->isEquality()) { + auto *tested = + getNonNullOperand(compare->getOperand(0), compare->getOperand(1)); + auto *other = getNonNullOperand(left, right); + bool same = + tested && other && + (tested == other || (tested == leftRoot && other == rightRoot) || + (tested == rightRoot && other == leftRoot)); + unsigned nullEdge = + compare->getPredicate() == llvm::CmpInst::ICMP_EQ ? 0 : 1; + if (same && branch->getSuccessor(nullEdge) == successor && + branch->getSuccessor(1 - nullEdge) != successor) + return true; + } + } + auto *leftPhi = llvm::dyn_cast(left); + auto *rightPhi = llvm::dyn_cast(right); + if (!leftPhi || !rightPhi || + leftPhi->getParent() != rightPhi->getParent()) + return false; + auto pair = std::make_pair(left, right); + if (llvm::is_contained(visited, pair)) + return true; + if (visited.size() == 16) + return false; + visited.push_back(pair); + for (unsigned index = 0; index < leftPhi->getNumIncomingValues(); + ++index) { + auto *block = leftPhi->getIncomingBlock(index); + if (!equivalent(leftPhi->getIncomingValue(index), + rightPhi->getIncomingValueForBlock(block), block, + leftPhi->getParent())) + return false; + } + return true; + }; + if (equivalent(leftRoot, rightRoot, nullptr, nullptr)) { + rightRoot->replaceAllUsesWith(leftRoot); + rightRoot->eraseFromParent(); + return true; + } + } + } + return false; + } + static llvm::Value *getNonNullOperand(llvm::Value *op1, llvm::Value *op2) { auto *ptr = llvm::dyn_cast(op1->getType()); if (!ptr) @@ -1126,6 +1198,42 @@ struct CoroBranchSimplifier : public llvm::PassInfoMixin { llvm::PreservedAnalyses run(llvm::Loop &loop, llvm::LoopAnalysisManager &am, llvm::LoopStandardAnalysisResults &ar, llvm::LPMUpdater &u) { + bool changed = mergeGuardedState(loop); + for (auto *block : loop.blocks()) { + for (auto &instruction : llvm::make_early_inc_range(*block)) { + auto *compare = llvm::dyn_cast(&instruction); + if (!compare || !compare->isEquality()) + continue; + auto *phi = llvm::dyn_cast_or_null( + getNonNullOperand(compare->getOperand(0), compare->getOperand(1))); + if (!phi || phi->getNumIncomingValues() > 8) + continue; + llvm::SmallVector values; + for (unsigned index = 0; index < phi->getNumIncomingValues(); ++index) { + auto *value = phi->getIncomingValue(index); + bool isNull = llvm::isa(value); + auto *predecessor = phi->getIncomingBlock(index); + llvm::SimplifyQuery query(phi->getModule()->getDataLayout(), &ar.DT, &ar.AC, + predecessor->getTerminator()); + if (!isNull && !llvm::isKnownNonZero(value, query)) + break; + values.push_back(llvm::ConstantInt::get( + compare->getType(), + isNull == (compare->getPredicate() == llvm::CmpInst::ICMP_EQ))); + } + if (values.size() != phi->getNumIncomingValues()) + continue; + auto *condition = llvm::PHINode::Create(compare->getType(), values.size(), + "coro.done", phi->getIterator()); + for (unsigned index = 0; index < values.size(); ++index) + condition->addIncoming(values[index], phi->getIncomingBlock(index)); + compare->replaceAllUsesWith(condition); + compare->eraseFromParent(); + changed = true; + } + } + if (changed) + return llvm::PreservedAnalyses::none(); if (auto *exit = loop.getExitingBlock()) { if (auto *br = llvm::dyn_cast(exit->getTerminator())) { if (!br->isConditional() || br->getNumSuccessors() != 2 || @@ -1180,6 +1288,89 @@ struct CoroBranchSimplifier : public llvm::PassInfoMixin { } }; +// Unrolling before coroutine lowering duplicates suspension states, which can leave +// a dispatch switch that prevents otherwise constant generator reductions from +// folding. Delay unrolling until those suspension points have been lowered away. +struct CoroUnrollControl : public llvm::PassInfoMixin { + static bool hasSuspension(llvm::BasicBlock &block) { + for (auto &instruction : block) + if (auto *intrinsic = llvm::dyn_cast(&instruction)) + if (intrinsic->getIntrinsicID() == llvm::Intrinsic::coro_suspend) + return true; + return false; + } + + static llvm::MDNode *updateMetadata(llvm::LLVMContext &context, + llvm::MDNode *identifier, bool suspends) { + auto nameOf = [](llvm::Metadata *metadata) -> llvm::StringRef { + auto *node = llvm::dyn_cast_or_null(metadata); + auto *name = node && node->getNumOperands() + ? llvm::dyn_cast_or_null(node->getOperand(0)) + : nullptr; + return name ? name->getString() : llvm::StringRef(); + }; + bool protectedLoop = false; + bool disabled = false; + if (identifier) + for (unsigned index = 1; index < identifier->getNumOperands(); ++index) { + auto name = nameOf(identifier->getOperand(index)); + protectedLoop |= name == "codon.coro.unroll"; + disabled |= name == "llvm.loop.unroll.disable"; + } + // Do not claim an existing disable hint: only our ownership marker authorizes + // removing it later. Repeated visits to an already protected loop are a no-op. + if ((suspends && disabled) || (!suspends && !protectedLoop)) + return identifier; + + // Preserve unrelated hints. Operand zero must refer to the new loop ID itself, + // so reserve it here and fill it after constructing the distinct metadata node. + llvm::SmallVector metadata{nullptr}; + if (identifier) + for (unsigned index = 1; index < identifier->getNumOperands(); ++index) { + auto name = nameOf(identifier->getOperand(index)); + if (name != "codon.coro.unroll" && name != "llvm.loop.unroll.disable") + metadata.push_back(identifier->getOperand(index)); + } + if (suspends) + for (auto name : {"codon.coro.unroll", "llvm.loop.unroll.disable"}) + metadata.push_back( + llvm::MDNode::get(context, llvm::MDString::get(context, name))); + auto *replacement = + metadata.size() > 1 ? llvm::MDNode::getDistinct(context, metadata) : nullptr; + if (replacement) + replacement->replaceOperandWith(0, replacement); + return replacement; + } + + llvm::PreservedAnalyses run(llvm::Loop &loop, llvm::LoopAnalysisManager &, + llvm::LoopStandardAnalysisResults &, llvm::LPMUpdater &) { + // Check this loop, not the whole function: non-suspending inner loops in a + // generator should retain ordinary unrolling even while its outer loop yields. + bool suspends = false; + for (auto *block : loop.blocks()) + suspends |= hasSuspension(*block); + loop.setLoopID( + updateMetadata(loop.getHeader()->getContext(), loop.getLoopID(), suspends)); + return llvm::PreservedAnalyses::all(); + } + + static void cleanup(llvm::Module &module) { + for (auto &function : module) { + if (llvm::any_of(function, hasSuspension)) + continue; + // Lowering can leave loop metadata on branches that are no longer latches. + // Once the function has no suspends, scan attachments directly so these + // stale guards cannot survive merely because LoopInfo no longer sees them. + for (auto &block : function) + for (auto &instruction : block) + if (auto *identifier = instruction.getMetadata(llvm::LLVMContext::MD_loop)) + instruction.setMetadata( + llvm::LLVMContext::MD_loop, + updateMetadata(module.getContext(), identifier, false)); + } + } +}; + struct OpenMPThreadIdOptimizer : public llvm::PassInfoMixin { llvm::PreservedAnalyses run(llvm::Function &function, llvm::FunctionAnalysisManager &) { @@ -1248,10 +1439,13 @@ struct OpenMPThreadIdOptimizer : public llvm::PassInfoMixindebug, options->jit); } diff --git a/codon/cir/transform/manager.cpp b/codon/cir/transform/manager.cpp index a4def4f50..49e0a062f 100644 --- a/codon/cir/transform/manager.cpp +++ b/codon/cir/transform/manager.cpp @@ -22,6 +22,7 @@ #include "codon/cir/transform/parallel/openmp.h" #include "codon/cir/transform/pass.h" #include "codon/cir/transform/pythonic/dict.h" +#include "codon/cir/transform/pythonic/enumerate.h" #include "codon/cir/transform/pythonic/format.h" #include "codon/cir/transform/pythonic/generator.h" #include "codon/cir/transform/pythonic/io.h" @@ -168,6 +169,7 @@ void PassManager::registerStandardPasses() { registerPass(std::make_unique()); registerPass(std::make_unique()); registerPass(std::make_unique()); + registerPass(std::make_unique()); registerPass(std::make_unique()); registerPass(std::make_unique()); @@ -210,6 +212,11 @@ void PassManager::registerStandardPasses() { registerPass(std::make_unique(numpyKey, seKey2), /*insertBefore=*/"", {numpyKey, seKey2}, {seKey1, rdKey, cfgKey, globalKey, capKey}); + // Expose whole producer/consumer loops before lowering. LLVM's suspension-aware + // unroll guard alone does not eliminate nested scan/consumer loop structure. + registerPass(std::make_unique(), + /*insertBefore=*/"", {}, + {seKey1, seKey2, rdKey, cfgKey, globalKey, capKey}); registerPass(std::make_unique(rdKey), /*insertBefore=*/"", {rdKey}, {seKey1, seKey2, rdKey, cfgKey, globalKey, capKey}); diff --git a/codon/cir/transform/pythonic/enumerate.cpp b/codon/cir/transform/pythonic/enumerate.cpp new file mode 100644 index 000000000..b2c3b7fc5 --- /dev/null +++ b/codon/cir/transform/pythonic/enumerate.cpp @@ -0,0 +1,178 @@ +// Copyright (C) 2022-2026 Exaloop Inc. + +#include "enumerate.h" + +#include "codon/cir/util/cloning.h" +#include "codon/cir/util/irtools.h" + +namespace codon { +namespace ir { +namespace transform { +namespace pythonic { +namespace { +// The yielded tuple can disappear only if its sole uses are the loop target +// definition and the two matched unpacking reads. +struct TupleUseChecker : public util::Operator { + ForFlow *loop; + Value *index; + Value *element; + bool valid = true; + + TupleUseChecker(ForFlow *loop, Value *index, Value *element) + : loop(loop), index(index), element(element) {} + + void preHook(Node *node) override { + auto *value = cast(node); + if (!value || value == loop || value == index || value == element) + return; + for (auto *var : value->getUsedVariables()) { + if (var->getId() == loop->getVar()->getId()) + valid = false; + } + } +}; + +AssignInstr *getUnpack(Value *value, Var *tuple, const std::string &field) { + auto *assign = cast(value); + auto *extract = assign ? cast(assign->getRhs()) : nullptr; + auto *var = extract ? util::getVar(extract->getVal()) : nullptr; + return var && var->getId() == tuple->getId() && extract->getField() == field + ? assign + : nullptr; +} +} // namespace + +const std::string EnumerateOptimization::KEY = "core-pythonic-enumerate-opt"; + +void EnumerateOptimization::handle(ForFlow *loop) { + // A shared incrementing counter assumes serial, synchronous iteration. + if (loop->isParallel() || loop->isAsync()) + return; + + // The local use check cannot see other functions observing a global target. + if (loop->getVar()->isGlobal()) + return; + + // Require a direct builtin call: stored iterators may have other consumers, + // and a user-defined enumerate need not have the builtin's semantics. + auto *call = cast(loop->getIter()); + auto *func = call ? util::getFunc(call->getCallee()) : nullptr; + if (!call || call->numArgs() != 2 || !func || + func->getName().rfind(ast::getMangledFunc("std.internal.builtin", "enumerate"), + 0) != 0) + return; + + // A loop target such as "index, value" becomes a temporary tuple variable + // followed by two leading assignments: index = tuple.item1, value = tuple.item2. + // Match these explicitly so we can replace the reads without constructing tuples. + auto *parent = cast(getParentFunc()); + auto *body = cast(loop->getBody()); + if (!parent || !body || body->begin() == body->end()) + return; + auto position = body->begin(); + auto *indexAssign = getUnpack(*position++, loop->getVar(), "item1"); + if (!indexAssign || position == body->end()) + return; + auto *elementAssign = getUnpack(*position++, loop->getVar(), "item2"); + if (!elementAssign) + return; + + // Check the whole function, not just the loop body: the tuple could also be + // read after the loop, in which case changing the loop variable would be unsafe. + TupleUseChecker uses(loop, cast(indexAssign->getRhs())->getVal(), + cast(elementAssign->getRhs())->getVal()); + uses.process(parent->getBody()); + if (!uses.valid) + return; + + auto *module = loop->getModule(); + auto *iterable = call->front(); + auto *start = call->back(); + if (!start->getType()->is(module->getIntType())) + return; + bool array = iterable->getType()->getName().rfind( + ast::getMangledClass("std.numpy.ndarray", "ndarray") + "[", 0) == 0; + Func *lenFunc = nullptr; + BodiedFunc *iterFunc = nullptr; + Value *iterExpr = nullptr; + // ndarrays support direct axis-zero indexing. Other iterables still need + // __iter__, unless the argument is already a generator. + if (array) { + lenFunc = module->getOrRealizeMethod(iterable->getType(), Module::LEN_MAGIC_NAME, + {iterable->getType()}); + if (!lenFunc) + return; + } else if (!isA(iterable->getType())) { + // Realize iter(x) through the frontend rather than calling the statically + // resolved __iter__ method: its return expression preserves virtual dispatch + // and default arguments. Decline if the helper has a more complicated body. + iterFunc = cast(module->getOrRealizeFunc("iter", {iterable->getType()}, + {}, "std.internal.builtin")); + if (!iterFunc) + return; + auto *iterBody = cast(iterFunc->getBody()); + if (!iterBody || iterBody->begin() == iterBody->end()) + return; + auto statement = iterBody->begin(); + auto *ret = cast(*statement++); + if (statement != iterBody->end() || !ret || !ret->getValue() || + !isA(ret->getValue()->getType())) + return; + iterExpr = ret->getValue(); + } + + // Evaluate source and start once, in that order, before iteration begins. + // Keep the counter private because the user's body may reassign its index target. + auto *setup = module->Nr(); + auto *source = util::makeVar(iterable, setup, parent); + auto *counter = util::makeVar(start, setup, parent); + auto *element = module->Nr(elementAssign->getLhs()->getType()); + parent->push_back(element); + indexAssign->setRhs(module->Nr(counter)); + elementAssign->setRhs(module->Nr(element)); + // Expose the current counter, then advance it before the user's body so that + // continue cannot skip the increment. + body->insert(position, + module->Nr(counter, *module->Nr(counter) + + *module->getInt(1))); + + if (array) { + // Use a zero-based offset independent of enumerate's start. __getitem__ + // preserves ndarray strides and yields row views for higher-dimensional arrays. + auto *offset = module->Nr(module->getIntType()); + parent->push_back(offset); + auto *length = util::call(lenFunc, {module->Nr(source)}); + auto *end = module->Nr(setup, length); + auto *load = (*module->Nr(source))[*module->Nr(offset)]; + // Fetch before assigning either visible target, just as a generator would: + // an indexing exception must leave their previous values intact. + body->insert(body->begin(), module->Nr(element, load)); + loop->replaceAll(module->N(loop->getSrcInfo(), module->getInt(0), + 1, end, body, offset)); + } else { + Value *iter = module->Nr(source); + if (iterExpr) { + util::CloneVisitor clone(module); + iter = + clone.clone(iterExpr, parent, {{(*iterFunc->arg_begin())->getId(), source}}); + } + // Preserve direct __iter__(source) calls so list lowering can recognize them. + // For virtual calls, setup must precede the entire expression, including the + // vtable lookup, rather than only evaluation of the receiver argument. + auto *iterCall = cast(iter); + if (util::isCallOf(iter, Module::ITER_MAGIC_NAME, 1) && + util::getVar(iterCall->front()) == source) { + iterCall->front()->replaceAll( + module->Nr(setup, module->Nr(source))); + } else { + iter = module->Nr(setup, iter); + } + loop->setIter(iter); + loop->setVar(element); + } +} + +} // namespace pythonic +} // namespace transform +} // namespace ir +} // namespace codon diff --git a/codon/cir/transform/pythonic/enumerate.h b/codon/cir/transform/pythonic/enumerate.h new file mode 100644 index 000000000..f04f07809 --- /dev/null +++ b/codon/cir/transform/pythonic/enumerate.h @@ -0,0 +1,22 @@ +// Copyright (C) 2022-2026 Exaloop Inc. + +#pragma once + +#include "codon/cir/transform/pass.h" + +namespace codon { +namespace ir { +namespace transform { +namespace pythonic { + +class EnumerateOptimization : public OperatorPass { +public: + static const std::string KEY; + std::string getKey() const override { return KEY; } + void handle(ForFlow *loop) override; +}; + +} // namespace pythonic +} // namespace transform +} // namespace ir +} // namespace codon diff --git a/codon/cir/transform/pythonic/generator.cpp b/codon/cir/transform/pythonic/generator.cpp index ae377c926..555c00b48 100644 --- a/codon/cir/transform/pythonic/generator.cpp +++ b/codon/cir/transform/pythonic/generator.cpp @@ -5,6 +5,7 @@ #include #include "codon/cir/util/cloning.h" +#include "codon/cir/util/inlining.h" #include "codon/cir/util/irtools.h" #include "codon/cir/util/matching.h" @@ -208,6 +209,236 @@ Func *genToAnyAll(BodiedFunc *gen, bool any) { } } // namespace +namespace { +struct FusionVerifier : public util::Operator { + int nodes = 0; + int yields = 0; + bool valid = true; + + void preHook(Node *) override { valid &= ++nodes <= 256; } + void handle(YieldInstr *value) override { + valid &= value->getValue() && !value->isFinal(); + ++yields; + } + void handle(YieldInInstr *) override { valid = false; } + void handle(AwaitInstr *) override { valid = false; } + void handle(TryCatchFlow *) override { valid = false; } + void handle(PointerValue *value) override { valid &= value->getVar()->isGlobal(); } + void handle(StackAllocInstr *) override { valid = false; } + void handle(ForFlow *loop) override { + valid &= !loop->isParallel() && !loop->isAsync(); + } +}; + +struct ConsumerVerifier : public util::Operator { + const std::unordered_set &wrappers; + bool valid = true; + explicit ConsumerVerifier(const std::unordered_set &wrappers) + : wrappers(wrappers) {} + void handle(BreakInstr *value) override { + valid &= value->getLoop() && wrappers.count(value->getLoop()->getId()); + } + void handle(ContinueInstr *) override { valid = false; } +}; + +struct IteratorUseVerifier : public util::Operator { + Var *iterator; + ForFlow *consumer; + const std::unordered_set &wrappers; + AssignInstr *assignment = nullptr; + int reads = 0; + int writes = 0; + int position = 0; + int creationPosition = 0; + int consumptionPosition = 0; + bool addressTaken = false; + std::vector creationLoops; + std::vector consumptionLoops; + IteratorUseVerifier(Var *iterator, ForFlow *consumer, + const std::unordered_set &wrappers) + : iterator(iterator), consumer(consumer), wrappers(wrappers) {} + std::vector enclosingLoops() { + std::vector result; + for (auto position = parent_begin(); position != parent_end(); ++position) { + auto *node = cast(*position); + if (node && node != consumer && !wrappers.count(node->getId()) && + (isA(node) || isA(node) || isA(node) || + isA(node) || isA(node))) { + result.push_back(node->getId()); + if ((isA(node) || isA(node)) && + position + 1 != parent_end()) + if (auto *branch = cast(*(position + 1))) + result.push_back(branch->getId()); + } + } + return result; + } + void preHook(Node *node) override { + ++position; + auto *value = cast(node); + if (!value || isA(value) || isA(value) || + isA(value)) + return; + for (auto *variable : value->getUsedVariables()) + addressTaken |= variable->getId() == iterator->getId(); + } + void handle(VarValue *value) override { + if (value->getVar()->getId() == iterator->getId()) { + ++reads; + consumptionPosition = position; + consumptionLoops = enclosingLoops(); + } + } + void handle(PointerValue *value) override { + addressTaken |= value->getVar()->getId() == iterator->getId(); + } + void handle(AssignInstr *value) override { + if (value->getLhs()->getId() == iterator->getId()) { + assignment = value; + ++writes; + creationPosition = position; + creationLoops = enclosingLoops(); + } + } + bool valid() const { + return reads == 1 && writes == 1 && !addressTaken && + creationPosition < consumptionPosition && creationLoops == consumptionLoops; + } +}; + +struct FusionTransformer : public util::Operator { + ForFlow *consumer; + WhileFlow *exit; + FusionTransformer(ForFlow *consumer, WhileFlow *exit) + : util::Operator(true), consumer(consumer), exit(exit) {} + + void handle(YieldInstr *value) override { + auto *module = value->getModule(); + value->replaceAll( + util::series(module->Nr(consumer->getVar(), value->getValue()), + consumer->getBody())); + } + void handle(ReturnInstr *value) override { + auto *module = value->getModule(); + auto *replacement = module->Nr(); + if (value->getValue()) + replacement->push_back(value->getValue()); + replacement->push_back(module->Nr(exit)); + value->replaceAll(replacement); + } +}; +} // namespace + +const std::string GeneratorLoopFusion::KEY = "core-pythonic-generator-loop-fusion"; + +void GeneratorLoopFusion::handle(CallInstr *call) { + auto *parent = cast(getParentFunc()); + auto *function = cast(util::getFunc(call->getCallee())); + if (!parent || !function || function->getName() != "__sum_wrapper") + return; + FusionVerifier verifier; + verifier.process(function->getBody()); + if (!verifier.valid || growth[parent->getId()] + verifier.nodes > 1024) + return; + auto inlined = util::inlineCall(call, /*aggressive=*/true); + if (!inlined) + return; + for (auto *variable : inlined.newVars) + parent->push_back(variable); + if (auto *expression = cast(inlined.result)) + if (auto *body = cast(expression->getFlow())) + if (auto *wrapper = cast(body->back())) + wrappers.insert(wrapper->getId()); + growth[parent->getId()] += verifier.nodes; + call->replaceAll(inlined.result); +} + +void GeneratorLoopFusion::handle(ForFlow *loop) { + if (loop->isParallel() || loop->isAsync()) + return; + auto *parent = cast(getParentFunc()); + if (!parent) + return; + if (auto *extract = cast(loop->getIter())) { + auto *variable = util::getVar(extract->getVal()); + if (!variable || variable->isGlobal()) + return; + IteratorUseVerifier uses(variable, loop, wrappers); + uses.process(parent->getBody()); + auto *tuple = uses.valid() ? cast(uses.assignment->getRhs()) : nullptr; + auto *constructor = tuple ? util::getFunc(tuple->getCallee()) : nullptr; + auto *type = constructor ? cast(constructor->getParentType()) : nullptr; + if (!type || type->getName() != "Tuple" || + constructor->getUnmangledName() != Module::NEW_MAGIC_NAME) + return; + auto index = cast(extract->getVal()->getType()) + ->getMemberIndex(extract->getField()); + if (index < 0 || index >= tuple->numArgs()) + return; + auto *setup = loop->getModule()->Nr(); + Var *selected = nullptr; + int position = 0; + for (auto *value : *tuple) { + auto *variable = util::makeVar(value, setup, parent); + if (position++ == index) + selected = variable; + } + uses.assignment->replaceAll(setup); + loop->setIter(loop->getModule()->Nr(selected)); + handle(loop); + return; + } + auto *call = cast(loop->getIter()); + AssignInstr *creation = nullptr; + if (auto *iterator = util::getVar(loop->getIter())) { + if (iterator->isGlobal()) + return; + IteratorUseVerifier uses(iterator, loop, wrappers); + uses.process(parent->getBody()); + if (uses.valid()) { + creation = uses.assignment; + call = cast(creation->getRhs()); + } + } + auto *generator = call ? cast(util::getFunc(call->getCallee())) : nullptr; + if (!generator || !generator->isGenerator() || !generator->getBody() || + generator->isAsync() || parent == generator || + call->numArgs() != std::distance(generator->arg_begin(), generator->arg_end()) || + util::hasAttribute(generator, + ast::getMangledFunc("std.internal.attributes", "noinline"))) + return; + FusionVerifier verifier; + verifier.process(generator->getBody()); + ConsumerVerifier consumerVerifier(wrappers); + consumerVerifier.process(loop->getBody()); + if (!verifier.valid || verifier.yields != 1 || !consumerVerifier.valid || + growth[parent->getId()] + verifier.nodes > 1024) + return; + growth[parent->getId()] += verifier.nodes; + + auto *module = loop->getModule(); + auto *setup = module->Nr(); + std::unordered_map arguments; + auto argument = generator->arg_begin(); + for (auto *value : *call) + arguments.emplace((*argument++)->getId(), util::makeVar(value, setup, parent)); + if (creation) { + creation->replaceAll(setup); + setup = module->Nr(); + } + util::CloneVisitor clone(module); + auto *body = cast(clone.clone(generator->getBody(), parent, arguments)); + auto *wrapper = module->Nr(); + auto *exit = module->Nr(module->getBool(true), wrapper); + wrappers.insert(exit->getId()); + FusionTransformer transformer(loop, exit); + transformer.process(body); + wrapper->push_back(body); + wrapper->push_back(module->Nr(exit)); + setup->push_back(exit); + loop->replaceAll(setup); +} + const std::string GeneratorArgumentOptimization::KEY = "core-pythonic-generator-argument-opt"; diff --git a/codon/cir/transform/pythonic/generator.h b/codon/cir/transform/pythonic/generator.h index 5ef8eab62..dad8c72ee 100644 --- a/codon/cir/transform/pythonic/generator.h +++ b/codon/cir/transform/pythonic/generator.h @@ -2,6 +2,9 @@ #pragma once +#include +#include + #include "codon/cir/transform/pass.h" namespace codon { @@ -19,6 +22,22 @@ class GeneratorArgumentOptimization : public OperatorPass { void handle(CallInstr *v) override; }; +class GeneratorLoopFusion : public OperatorPass { + std::unordered_set wrappers; + std::unordered_map growth; + +public: + static const std::string KEY; + std::string getKey() const override { return KEY; } + void run(Module *module) override { + wrappers.clear(); + growth.clear(); + OperatorPass::run(module); + } + void handle(CallInstr *call) override; + void handle(ForFlow *loop) override; +}; + } // namespace pythonic } // namespace transform } // namespace ir diff --git a/stdlib/algorithms/timsort.codon b/stdlib/algorithms/timsort.codon index 77c21fe5e..7d9a6d43e 100644 --- a/stdlib/algorithms/timsort.codon +++ b/stdlib/algorithms/timsort.codon @@ -365,6 +365,14 @@ def _tim_sort( if end - begin < 2: return + if end - begin < 64: + run_length, inorder = _count_run(arr, begin, end, keyf) + if not inorder: + _reverse_sortslice(arr, begin, begin + run_length) + if run_length < end - begin: + _insertion_sort(arr, begin, end, keyf) + return + merge_pending = List[Tuple[int, int]]() minrun = _merge_compute_minrun(end - begin) i = begin diff --git a/stdlib/heapq.codon b/stdlib/heapq.codon index 37325005d..297b9fbd3 100644 --- a/stdlib/heapq.codon +++ b/stdlib/heapq.codon @@ -150,12 +150,12 @@ def nsmallest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> L result = List(n) done = False for i in range(n): - if it.done(): + if not it._advance(): done = True break - result.append((it.__next__(), i)) + result.append((it._get(), i)) if not result: - it.destroy() + it._destroy() return [] _heapify_max(result) top = result[0][0] @@ -167,7 +167,7 @@ def nsmallest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> L top, _order = result[0] order += 1 else: - it.destroy() + it._destroy() result.sort() return [elem for elem, order in result] else: @@ -176,13 +176,13 @@ def nsmallest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> L result = List(n) done = False for i in range(n): - if it.done(): + if not it._advance(): done = True break - elem = it.__next__() + elem = it._get() result.append((key(elem), i, elem)) if not result: - it.destroy() + it._destroy() return [] _heapify_max(result) top = result[0][0] @@ -195,7 +195,7 @@ def nsmallest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> L top, _order, _elem = result[0] order += 1 else: - it.destroy() + it._destroy() result.sort() return [elem for k, order, elem in result] @@ -219,12 +219,12 @@ def nlargest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> Li result = List(n) done = False for i in range(0, -n, -1): - if it.done(): + if not it._advance(): done = True break - result.append((it.__next__(), i)) + result.append((it._get(), i)) if not result: - it.destroy() + it._destroy() return [] heapify(result) top = result[0][0] @@ -236,7 +236,7 @@ def nlargest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> Li top, _order = result[0] order -= 1 else: - it.destroy() + it._destroy() result.sort() return [elem for elem, order in reversed(result)] else: @@ -245,10 +245,10 @@ def nlargest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> Li result = List(n) done = False for i in range(0, -n, -1): - if it.done(): + if not it._advance(): done = True break - elem = it.__next__() + elem = it._get() result.append((key(elem), i, elem)) if not result: return [] @@ -263,7 +263,7 @@ def nlargest(n: int, iterable: Generator[T], key=Optional[int](), T: type) -> Li top, _order, _elem = result[0] order -= 1 else: - it.destroy() + it._destroy() result.sort() return [elem for k, order, elem in reversed(result)] @@ -304,8 +304,8 @@ def merge(*iterables, key=Optional[int](), reverse: bool = False): order = 0 for it in iterables: gen = iter(it) - if not gen.done(): - items.append(_MergeItem(gen.__next__(), order * direction, gen, key)) + if gen._advance(): + items.append(_MergeItem(gen._get(), order * direction, gen, key)) order += 1 _heapify(items) while len(items) > 1: @@ -313,10 +313,10 @@ def merge(*iterables, key=Optional[int](), reverse: bool = False): # TODO: @tuple unpacking does not work value, order, gen = items[0].value, items[0].order, items[0].gen yield value - if gen.done(): + if not gen._advance(): _heappop(items) break - _heapreplace(items, _MergeItem(gen.__next__(), order, gen, key)) + _heapreplace(items, _MergeItem(gen._get(), order, gen, key)) if items: # fast case when only a single iterator remains value, order, gen = items[0].value, items[0].order, items[0].gen diff --git a/stdlib/internal/__init_test__.codon b/stdlib/internal/__init_test__.codon index 7c3159699..b3ea9d40c 100644 --- a/stdlib/internal/__init_test__.codon +++ b/stdlib/internal/__init_test__.codon @@ -119,12 +119,12 @@ from internal.types.str import * from internal.format import * def next(g: Generator[T], default: Optional[T] = None, T: type) -> T: - if g.done(): - if default: - return unwrap(default) + if not g._advance(): + if default is not None: + return default.__val__() else: raise StopIteration() - return g.__next__() + return g._get() from C import seq_print_full(str, cobj) diff --git a/stdlib/internal/builtin.codon b/stdlib/internal/builtin.codon index d7160dce1..50080ba15 100644 --- a/stdlib/internal/builtin.codon +++ b/stdlib/internal/builtin.codon @@ -72,20 +72,20 @@ def min(*args, key=None, default=None): compile_error("min() 'default' argument only allowed for iterables") elif static.len(args) == 1: x = args[0].__iter__() - if not x.done(): - s = x.__next__() - while not x.done(): - i = x.__next__() + if x._advance(): + s = x._get() + while x._advance(): + i = x._get() if key is None: if i < s: s = i else: if key(i) < key(s): s = i - x.destroy() + x._destroy() return s else: - x.destroy() + x._destroy() if default is None: raise ValueError("min() arg is an empty sequence") else: @@ -114,20 +114,20 @@ def max(*args, key=None, default=None): compile_error("max() 'default' argument only allowed for iterables") elif static.len(args) == 1: x = args[0].__iter__() - if not x.done(): - s = x.__next__() - while not x.done(): - i = x.__next__() + if x._advance(): + s = x._get() + while x._advance(): + i = x._get() if key is None: if i > s: s = i else: if key(i) > key(s): s = i - x.destroy() + x._destroy() return s else: - x.destroy() + x._destroy() if default is None: raise ValueError("max() arg is an empty sequence") else: @@ -194,12 +194,12 @@ def chr(codepoint: int) -> str: return str(data, 1 | (kind << 56)) def next(g: Generator[T], default: Optional[T] = None, T: type) -> T: - if g.done(): + if not g._advance(): if default is not None: return default.__val__() else: raise StopIteration() - return g.__next__() + return g._get() def any(x: Generator[T], T: type) -> bool: for a in x: @@ -218,15 +218,28 @@ def zip(*args): yield from List[int]() else: iters = tuple(iter(i) for i in args) - done = False - while not done: - for i in iters: - if i.done(): - done = True - if not done: - yield tuple(i.__next__() for i in iters) - for i in iters: - i.destroy() + + @pure + @llvm + def empty_result(R: type) -> R: + ret {=R} zeroinitializer + + @pure + @derives + @llvm + def set_result(result: R, value: V, index: Literal[int], R: type, V: type) -> R: + %updated = insertvalue {=R} %result, {=V} %value, {=index} + ret {=R} %updated + + Result = type(tuple(iterator._get() for iterator in iters)) + while True: + result = empty_result(Result) + for index in static.range(static.len(iters)): + iterator = iters[index] + if not iterator._advance(): + return + result = set_result(result, iterator._get(), index) + yield result def filter(f: CallableTrait[[T], bool], x: Generator[T], T: type) -> Generator[T]: for a in x: @@ -427,6 +440,13 @@ _CHAR_POS: Literal[int] = 43 _CHAR_NEG: Literal[int] = 45 def _parse_int(s: str, base: int = 10, T: type = int): + if s._width() == 1: + return _parse_int_typed(s, s._ptr, base, T) + if s._width() == 2: + return _parse_int_typed(s, Ptr[u16](s._ptr), base, T) + return _parse_int_typed(s, Ptr[u32](s._ptr), base, T) + +def _parse_int_typed(s: str, data: Ptr[Char], base: int, T: type, Char: type): def _digit_value(c: int): if _CHAR_0 <= c and c <= _CHAR_9: @@ -474,15 +494,15 @@ def _parse_int(s: str, base: int = 10, T: type = int): end = len(s) # skip leading whitespace - while p < end and _isspace(s._load_codepoint(p)): + while p < end and _isspace(int(data[p])): p += 1 # sign neg = False if p < end and ( - s._load_codepoint(p) == _CHAR_POS or s._load_codepoint(p) == _CHAR_NEG + int(data[p]) == _CHAR_POS or int(data[p]) == _CHAR_NEG ): - neg = s._load_codepoint(p) == _CHAR_NEG + neg = int(data[p]) == _CHAR_NEG p += 1 if not _signed(T): @@ -493,42 +513,31 @@ def _parse_int(s: str, base: int = 10, T: type = int): base0_leading_zero = False if base == 0: base = 10 - if (end - p) >= 2 and s._load_codepoint(p) == _CHAR_0: - if s._load_codepoint(p + 1) == _CHAR_x or s._load_codepoint(p + 1) == _CHAR_X: + if (end - p) >= 2 and int(data[p]) == _CHAR_0: + if int(data[p + 1]) == _CHAR_x or int(data[p + 1]) == _CHAR_X: base = 16 p += 2 - if p < end and s._load_codepoint(p) == _CHAR_UNDERSCORE: + if p < end and int(data[p]) == _CHAR_UNDERSCORE: p += 1 - elif s._load_codepoint(p + 1) == _CHAR_b or s._load_codepoint(p + 1) == _CHAR_B: + elif int(data[p + 1]) == _CHAR_b or int(data[p + 1]) == _CHAR_B: base = 2 p += 2 - if p < end and s._load_codepoint(p) == _CHAR_UNDERSCORE: + if p < end and int(data[p]) == _CHAR_UNDERSCORE: p += 1 - elif s._load_codepoint(p + 1) == _CHAR_o or s._load_codepoint(p + 1) == _CHAR_O: + elif int(data[p + 1]) == _CHAR_o or int(data[p + 1]) == _CHAR_O: base = 8 p += 2 - if p < end and s._load_codepoint(p) == _CHAR_UNDERSCORE: + if p < end and int(data[p]) == _CHAR_UNDERSCORE: p += 1 else: base0_leading_zero = True - else: - if (base == 16 and (end - p) >= 2 and - s._load_codepoint(p) == _CHAR_0 and - (s._load_codepoint(p + 1) == _CHAR_x or s._load_codepoint(p + 1) == _CHAR_X)): - p += 2 - if p < end and s._load_codepoint(p) == _CHAR_UNDERSCORE: - p += 1 - elif (base == 8 and (end - p) >= 2 and - s._load_codepoint(p) == _CHAR_0 and - (s._load_codepoint(p + 1) == _CHAR_o or s._load_codepoint(p + 1) == _CHAR_O)): + elif (end - p) >= 2 and int(data[p]) == _CHAR_0: + prefix = int(data[p + 1]) + if ((base == 16 and (prefix == _CHAR_x or prefix == _CHAR_X)) or + (base == 8 and (prefix == _CHAR_o or prefix == _CHAR_O)) or + (base == 2 and (prefix == _CHAR_b or prefix == _CHAR_B))): p += 2 - if p < end and s._load_codepoint(p) == _CHAR_UNDERSCORE: - p += 1 - elif (base == 2 and (end - p) >= 2 and - s._load_codepoint(p) == _CHAR_0 and - (s._load_codepoint(p + 1) == _CHAR_b or s._load_codepoint(p + 1) == _CHAR_B)): - p += 2 - if p < end and s._load_codepoint(p) == _CHAR_UNDERSCORE: + if p < end and int(data[p]) == _CHAR_UNDERSCORE: p += 1 limit = _limit(neg, T) @@ -552,7 +561,7 @@ def _parse_int(s: str, base: int = 10, T: type = int): raise ValueError(f"invalid literal for int() with base {base}: {s.__repr__()}") while p < end: - c = s._load_codepoint(p) + c = int(data[p]) if c == _CHAR_UNDERSCORE: if (not saw_digit) or prev_underscore: @@ -588,7 +597,7 @@ def _parse_int(s: str, base: int = 10, T: type = int): invalid(s, base) # skip trailing whitespace - while p < end and _isspace(s._load_codepoint(p)): + while p < end and _isspace(int(data[p])): p += 1 if p != end: @@ -620,6 +629,14 @@ class UInt: return _parse_int(s, base, UInt[N]) def _parse_float_prefix(value: str, start: int): + if start == len(value): + return 0.0, 0 + if value._is_ascii(): + data = value._ptr + start + end = cobj() + result = _C.seq_float_from_str(str(data, len(value) - start), __ptr__(end)) + return result, end - data + length = 0 data = Ptr[u8](len(value) - start) diff --git a/stdlib/internal/python.codon b/stdlib/internal/python.codon index e5803a43a..bd9a02707 100644 --- a/stdlib/internal/python.codon +++ b/stdlib/internal/python.codon @@ -2281,11 +2281,11 @@ class _PyWrap: return cobj() gt = TO(pt._gen) - if gt.done(): + if not gt._advance(): pt._gen = cobj() return cobj() else: - return gt.__next__().__to_py__() + return gt._get().__to_py__() def __to_py__(self): return _PyWrap.wrap_to_py(self) diff --git a/stdlib/internal/types/collections/list.codon b/stdlib/internal/types/collections/list.codon index 437a3beab..c9f13f563 100644 --- a/stdlib/internal/types/collections/list.codon +++ b/stdlib/internal/types/collections/list.codon @@ -244,19 +244,21 @@ class List: return self def __mul__(self, n: int) -> List[T]: - if n <= 0: + if n <= 0 or self._len == 0: return List[T]() new_len = self._len * n - v = List[T](new_len) - i = 0 - while i < n: - j = 0 - while j < self._len: - v.append(self._get(j)) - j += 1 - i += 1 - return v + data = Ptr[T](new_len) + if self._len == 1: + value = self._get(0) + for index in range(n): + data[index] = value + else: + block_size = self._len * gc.sizeof(T) + for index in range(n): + str.memcpy((data + index * self._len).as_byte(), + self._ptr.as_byte(), block_size) + return List[T](data, new_len) def __rmul__(self, n: int) -> List[T]: return self.__mul__(n) diff --git a/stdlib/internal/types/generator.codon b/stdlib/internal/types/generator.codon index 43ee1d7ac..b5c1d5352 100644 --- a/stdlib/internal/types/generator.codon +++ b/stdlib/internal/types/generator.codon @@ -14,16 +14,25 @@ class Generator: def __promise__(self) -> Ptr[T]: pass - def done(self) -> bool: + def _advance(self) -> bool: + """Resume once unless exhausted; return whether a value was yielded.""" + if self.__done__(): + return False self.__resume__() - return self.__done__() + return not self.__done__() - def __next__(self: Generator[T]) -> T: + def _get(self: Generator[T]) -> T: + """Read the current value after a successful `_advance()`, without resuming.""" if isinstance(T, None): pass else: return self.__promise__()[0] + def __next__(self: Generator[T]) -> T: + if not self._advance(): + raise StopIteration() + return self._get() + def __iter__(self) -> Generator[T]: return self @@ -62,14 +71,15 @@ class Generator: return Ptr._ptr_to_str(self.__raw__(), "generator") def send(self, what: T) -> T: - p = self.__promise__() - p[0] = what - self.__resume__() - return p[0] + if self.__done__(): + raise StopIteration() + if not isinstance(T, None): + self.__promise__()[0] = what + return self.__next__() @nocapture @llvm - def destroy(self) -> None: + def _destroy(self) -> None: declare void @llvm.coro.destroy(ptr) call void @llvm.coro.destroy(ptr %self) ret {} {} diff --git a/stdlib/internal/types/str.codon b/stdlib/internal/types/str.codon index 5258ca8a3..951588698 100644 --- a/stdlib/internal/types/str.codon +++ b/stdlib/internal/types/str.codon @@ -5099,6 +5099,18 @@ def _unicode_join_impl(sep: str, values: List[str]) -> str: result = _unicode_new_uninit(total_len, maxchar) out = 0 + if result._kind() <= KIND_LATIN1: + for index in range(len(values)): + if index > 0 and len(sep) > 0: + str.memcpy(result._ptr + out, sep._ptr, len(sep)) + out += len(sep) + + piece = values[index] + str.memcpy(result._ptr + out, piece._ptr, len(piece)) + out += len(piece) + + return result + for i in range(len(values)): if i > 0 and len(sep) > 0: _unicode_copy_range(result, sep, out, 0, len(sep)) diff --git a/stdlib/itertools.codon b/stdlib/itertools.codon index 687eabbb4..882c2883b 100644 --- a/stdlib/itertools.codon +++ b/stdlib/itertools.codon @@ -54,8 +54,8 @@ class chain: @inline def __new__(*iterables): - for it in iterables: - for element in it: + for index in static.range(static.len(iterables)): + for element in iterables[index]: yield element @inline @@ -91,14 +91,13 @@ def filterfalse( if not predicate(x): yield x -# TODO: fix key @inline -def groupby(iterable, key=Optional[int]()): +def groupby(iterable, key=None): currkey = None group = [] for currvalue in iterable: - k = currvalue if isinstance(key, Optional) else key(currvalue) + k = currvalue if key is None else key(currvalue) if currkey is None: currkey = k if k != unwrap(currkey): @@ -114,12 +113,14 @@ def islice(iterable: Generator[T], stop: Optional[int], T: type) -> Generator[T] raise ValueError( "Indices for islice() must be None or an integer: 0 <= x <= sys.maxsize." ) - i = 0 - for x in iterable: - if stop is not None and i >= stop.__val__(): - break - yield x - i += 1 + if stop is not None and stop.__val__() == 0: + return + index = 0 + for value in iterable: + yield value + index += 1 + if stop is not None and index >= stop.__val__(): + return @overload def islice( @@ -132,40 +133,27 @@ def islice( from sys import maxsize start: int = 0 if start is None else start - stop: int = maxsize if stop is None else stop step: int = 1 if step is None else step - have_stop = False - if start < 0 or stop < 0: + if start < 0 or (stop is not None and stop.__val__() < 0): raise ValueError( "Indices for islice() must be None or an integer: 0 <= x <= sys.maxsize." ) - elif step < 0: + elif step <= 0: raise ValueError("Step for islice() must be a positive integer or None.") - it = range(start, stop, step) - N = len(it) - idx = 0 - b = -1 - - if N == 0: - for i, element in zip(range(start), iterable): - pass + limit = 0 if stop is None else max(start, stop.__val__()) + if stop is not None and limit == 0: return - - nexti = it[0] - for i, element in enumerate(iterable): - if i == nexti: - yield element - idx += 1 - if idx >= N: - b = i - break - nexti = it[idx] - - if b >= 0: - for i, element in zip(range(b + 1, stop), iterable): - pass + index = 0 + next_index = start + for value in iterable: + if index == next_index: + yield value + next_index = next_index + step if step <= maxsize - next_index else -1 + index += 1 + if stop is not None and index >= limit: + return @inline def starmap(function, iterable): @@ -191,12 +179,9 @@ def tee(iterable: Generator[T], n: int = 2, T: type) -> List[Generator[T]]: def gen(mydeque: deque[T], T: type) -> Generator[T]: while True: if not mydeque: # when the local deque is empty - if it.__done__(): - return - it.__resume__() - if it.__done__(): + if not it._advance(): return - newval = it.__next__() + newval = it._get() for d in deques: # load it to all the deques d.append(newval) yield mydeque.popleft() @@ -211,21 +196,21 @@ def zip_longest(*iterables, fillvalue): a_done = False b_done = False - while not a.done(): - a_val = a.__next__() + while a._advance(): + a_val = a._get() b_val = fillvalue if not b_done: - b_done = b.done() + b_done = not b._advance() if not b_done: - b_val = b.__next__() + b_val = b._get() yield a_val, b_val if not b_done: - while not b.done(): - yield fillvalue, b.__next__() + while b._advance(): + yield fillvalue, b._get() - a.destroy() - b.destroy() + a._destroy() + b._destroy() else: iterators = tuple(iter(it) for it in iterables) num_active = len(iterators) @@ -236,13 +221,13 @@ def zip_longest(*iterables, fillvalue): for it in iterators: if it.__done__(): # already done values.append(fillvalue) - elif it.done(): # resume and check + elif not it._advance(): num_active -= 1 if not num_active: return values.append(fillvalue) else: - values.append(it.__next__()) + values.append(it._get()) yield values @inline @@ -250,9 +235,9 @@ def zip_longest(*iterables, fillvalue): def zip_longest(*args): def get_next(it): - if it.__done__() or it.done(): + if not it._advance(): return None - return it.__next__() + return it._get() iters = tuple(iter(arg) for arg in args) while True: @@ -266,7 +251,7 @@ def zip_longest(*args): return yield result for it in iters: - it.destroy() + it._destroy() # Combinatoric iterators diff --git a/stdlib/math.codon b/stdlib/math.codon index 24f8f0467..e09c64864 100644 --- a/stdlib/math.codon +++ b/stdlib/math.codon @@ -19,32 +19,17 @@ inf = _inf() nan = _nan() def factorial(x: int) -> int: - _F = ( - 1, - 1, - 2, - 6, - 24, - 120, - 720, - 5040, - 40320, - 362880, - 3628800, - 39916800, - 479001600, - 6227020800, - 87178291200, - 1307674368000, - 20922789888000, - 355687428096000, - 6402373705728000, - 121645100408832000, - 2432902008176640000, - ) + @pure + @llvm + def lookup(index: int) -> int: + @data = private unnamed_addr constant [21 x i64] [i64 1, i64 1, i64 2, i64 6, i64 24, i64 120, i64 720, i64 5040, i64 40320, i64 362880, i64 3628800, i64 39916800, i64 479001600, i64 6227020800, i64 87178291200, i64 1307674368000, i64 20922789888000, i64 355687428096000, i64 6402373705728000, i64 121645100408832000, i64 2432902008176640000], align 8 + %pointer = getelementptr inbounds [21 x i64], ptr @data, i64 0, i64 %index + %value = load i64, ptr %pointer, align 8, !range !{i64 1, i64 2432902008176640001} + ret i64 %value + if not (0 <= x <= 20): raise ValueError("factorial is only supported for 0 <= x <= 20") - return _F[x] + return lookup(x) def isnan(x: float) -> bool: @pure diff --git a/stdlib/re.codon b/stdlib/re.codon index 9936e30b8..ccb11a343 100644 --- a/stdlib/re.codon +++ b/stdlib/re.codon @@ -301,6 +301,9 @@ class Match[T]: return r.__str__() def expand(self, template: str): + if "\\" not in template: + return template + def get_or_empty(s: Optional[str]): return s if s is not None else '' @@ -434,10 +437,7 @@ class Pattern: return spans def _raw_spans(self, raw_spans: Ptr[Span]): - spans = Ptr[Span](self.groups + 1) - for i in range(self.groups + 1): - spans[i] = Span(raw_spans[i].start, raw_spans[i].end) - return spans + return raw_spans def _match_one( self, anchor: int, string: T, pos: Optional[int], endpos: Optional[int] @@ -449,9 +449,10 @@ class Pattern: return None encoded = string.encode() if isinstance(string, str) else string - byte_pos = _utf8_byte_offset(encoded, posx) if isinstance(string, str) else posx + unicode_offsets = len(encoded) != len(string) + byte_pos = _utf8_byte_offset(encoded, posx) if unicode_offsets else posx byte_endpos = len(encoded) if endposx == len(string) else ( - _utf8_byte_offset(encoded, endposx) if isinstance(string, str) else endposx + _utf8_byte_offset(encoded, endposx) if unicode_offsets else endposx ) raw_spans = seq_re_match(self._re, anchor, encoded, byte_pos, byte_endpos) @@ -460,7 +461,7 @@ class Pattern: spans = ( self._match_spans(raw_spans, encoded, byte_pos, posx) - if isinstance(string, str) else self._raw_spans(raw_spans) + if unicode_offsets else self._raw_spans(raw_spans) ) return Match[T](spans, posx, endposx, self, string) @@ -474,9 +475,10 @@ class Pattern: return encoded = string.encode() if isinstance(string, str) else string - byte_pos = _utf8_byte_offset(encoded, posx) if isinstance(string, str) else posx + unicode_offsets = len(encoded) != len(string) + byte_pos = _utf8_byte_offset(encoded, posx) if unicode_offsets else posx byte_endpos = len(encoded) if endposx == len(string) else ( - _utf8_byte_offset(encoded, endposx) if isinstance(string, str) else endposx + _utf8_byte_offset(encoded, endposx) if unicode_offsets else endposx ) codepoint_pos = posx @@ -488,7 +490,7 @@ class Pattern: spans = ( self._match_spans(raw_spans, encoded, byte_pos, codepoint_pos) - if isinstance(string, str) else self._raw_spans(raw_spans) + if unicode_offsets else self._raw_spans(raw_spans) ) yield Match[T](spans, posx, endposx, self, string) @@ -501,7 +503,7 @@ class Pattern: byte_pos = raw_spans[0].end posx = spans[0].end codepoint_pos = posx - if isinstance(string, str): + if unicode_offsets: c = int(encoded._ptr[byte_pos]) byte_pos += 1 if c < 0x80 else 2 if c < 0xE0 else 3 if c < 0xF0 else 4 else: @@ -580,7 +582,7 @@ class Pattern: return repl(match) def subn(self, repl, string: T, count: int = 0): - cb = lambda match: [Pattern._repl(match, repl)] + cb = lambda match: (Pattern._repl(match, repl),) pieces, numsplit = self._split(cb, string, count) joined_pieces = str.cat(pieces) if isinstance(string, str) else bytes.cat(pieces) return joined_pieces, numsplit diff --git a/test/cir/llvm/optimize.cpp b/test/cir/llvm/optimize.cpp index f35c42ae1..927d8b380 100644 --- a/test/cir/llvm/optimize.cpp +++ b/test/cir/llvm/optimize.cpp @@ -93,10 +93,13 @@ int countLazyFixedAllocationCaches(llvm::Module *module, uint64_t size) { return count; } -std::unique_ptr compileAndOptimize(const std::string &code) { +std::unique_ptr +compileAndOptimize(const std::string &code, + const std::vector &disabled = {}) { auto options = Options::getDefault("build/codon_test"); options->debug = false; options->standalone = true; + options->disabled.insert(options->disabled.end(), disabled.begin(), disabled.end()); auto compiler = std::make_unique(*options); llvm::cantFail(compiler->parseCode("allocation_phi_test.codon", code)); llvm::cantFail(compiler->compile()); @@ -306,6 +309,609 @@ TEST(LLVMOptimizationTest, RemovesUnusedStandardStreamInitialization) { EXPECT_EQ(1, definitions); } +TEST(LLVMOptimizationTest, ElidesSafeGeneratorAndZipIteration) { + ASSERT_EXIT( + { + auto compiler = compileAndOptimize(R"( +def generator_values(data: Ptr[int], count: int): + index = 0 + while index < count: + yield data[index] + index += 1 + +@export +def consume_generator(data: Ptr[int], count: int) -> int: + total = 0 + for value in generator_values(data, count): + total += value + return total + +@export +def consume_zip(data: Ptr[int], count: int) -> int: + total = 0 + for left, right in zip(generator_values(data, count), generator_values(data, count)): + total += left + right + return total + +@export +def consume_aliased_zip(data: Ptr[int], count: int) -> int: + total = 0 + items = generator_values(data, count) + for left, right in zip(items, items): + total += left + right + return total + +def range_values(data: Ptr[int], count: int): + for index in range(count): + yield data[index] + +def delegated_values(data: Ptr[int], count: int): + yield from range_values(data, count) + yield from range_values(data + 1, count) + +@export +def consume_delegated(data: Ptr[int], count: int) -> int: + total = 0 + for value in delegated_values(data, count): + total += value + return total + +@export +def consume_nested_zip(data: Ptr[int], count: int) -> int: + total = 0 + inner = zip(range_values(data, count), range_values(data + 1, count)) + for pair, third in zip(inner, range_values(data + 2, count)): + total += pair[0] * pair[1] + third + return total + +@export +def escaping_generator(data: Ptr[int], count: int): + return generator_values(data, count) +)"); + auto *module = compiler->getLLVMVisitor()->getModule(); + EXPECT_FALSE(llvm::verifyModule(*module, &llvm::errs())); + for (const auto &name : + {"consume_generator", "consume_zip", "consume_aliased_zip", + "consume_delegated", "consume_nested_zip"}) { + SCOPED_TRACE(name); + auto *function = module->getFunction(name); + ASSERT_NE(nullptr, function); + for (auto &block : *function) { + for (auto &instruction : block) { + if (auto *call = llvm::dyn_cast(&instruction)) { + auto *callee = call->getCalledFunction(); + EXPECT_TRUE(callee && callee->isIntrinsic() && + !callee->getName().starts_with("llvm.coro.")) + << "generator iteration retained a call in " << name; + } + } + } + } + for (const auto *name : {"consume_delegated", "consume_nested_zip"}) { + SCOPED_TRACE(name); + bool vectorized = false; + for (auto &block : *module->getFunction(name)) + for (auto &instruction : block) + vectorized |= llvm::isa(instruction) && + instruction.getType()->isVectorTy(); + EXPECT_TRUE(vectorized); + } + auto *escaping = module->getFunction("escaping_generator"); + ASSERT_NE(nullptr, escaping); + bool allocates = false; + for (auto &block : *escaping) { + for (auto &instruction : block) { + if (auto *call = llvm::dyn_cast(&instruction)) { + auto *callee = call->getCalledFunction(); + allocates |= callee && callee->getName() == "seq_alloc"; + } + } + } + EXPECT_TRUE(allocates); + std::_Exit(HasFailure() ? EXIT_FAILURE : EXIT_SUCCESS); + }, + testing::ExitedWithCode(EXIT_SUCCESS), ""); +} + +TEST(LLVMOptimizationTest, ElidesChainedAndSlicedGenerators) { + ASSERT_EXIT( + { + auto compiler = compileAndOptimize(R"( +from itertools import chain, islice + +def slice_values(data: Ptr[int], count: int): + for index in range(count): + yield data[index] + +@export +def chained_sum(data: Ptr[int], count: int) -> int: + return sum(chain(slice_values(data, count), slice_values(data + 1, count))) + +@export +def chained_loop(data: Ptr[int], count: int) -> int: + total = 0 + for value in chain(slice_values(data, count), slice_values(data + 1, count)): + total += value + return total + +@export +def prefix_sum(data: Ptr[int], count: int) -> int: + return sum(islice(slice_values(data, count), max(count // 2, 0))) + +@export +def prefix_loop(data: Ptr[int], count: int) -> int: + total = 0 + for value in islice(slice_values(data, count), max(count // 2, 0)): + total += value + return total + +@export +def prefix_manual(data: Ptr[int], count: int) -> int: + total = 0 + for index in range(max(count // 2, 0)): + total += data[index] + return total + +@export +def strided_sum(data: Ptr[int], count: int) -> int: + return sum(islice(slice_values(data, count), 1, max(count, 0), 2)) + +@export +def strided_loop(data: Ptr[int], count: int) -> int: + total = 0 + for value in islice(slice_values(data, count), 1, max(count, 0), 2): + total += value + return total + +@export +def sliced_chain_loop(data: Ptr[int], count: int) -> int: + total = 0 + items = chain(slice_values(data, count), slice_values(data + 1, count)) + for value in islice(items, 1, max(count, 0)): + total += value + return total +)"); + auto *module = compiler->getLLVMVisitor()->getModule(); + EXPECT_FALSE(llvm::verifyModule(*module, &llvm::errs())); + for (auto *name : {"chained_sum", "chained_loop", "prefix_sum", "prefix_loop", + "strided_sum", "strided_loop", "sliced_chain_loop"}) { + SCOPED_TRACE(name); + auto *function = module->getFunction(name); + ASSERT_NE(nullptr, function); + for (auto &block : *function) + for (auto &instruction : block) { + EXPECT_FALSE(llvm::isa(instruction)); + if (auto *call = llvm::dyn_cast(&instruction)) { + auto *callee = call->getCalledFunction(); + EXPECT_TRUE(callee && callee->isIntrinsic() && + !callee->getName().starts_with("llvm.coro.")) + << (callee ? callee->getName().str() : "indirect call"); + } + } + } + auto vectorAccumulators = [](llvm::Function *function) { + unsigned maximum = 0; + for (auto &block : *function) { + unsigned count = 0; + for (auto &instruction : block) + count += llvm::isa(instruction) && + instruction.getType()->isVectorTy(); + maximum = std::max(maximum, count); + } + return maximum; + }; + auto expected = vectorAccumulators(module->getFunction("prefix_manual")); + EXPECT_GT(expected, 0u); + EXPECT_EQ(expected, vectorAccumulators(module->getFunction("prefix_sum"))); + EXPECT_EQ(expected, vectorAccumulators(module->getFunction("prefix_loop"))); + std::_Exit(HasFailure() ? EXIT_FAILURE : EXIT_SUCCESS); + }, + testing::ExitedWithCode(EXIT_SUCCESS), ""); +} + +TEST(LLVMOptimizationTest, PreservesUserUnrollMetadata) { + // Without our ownership marker, the disable hint and unrelated metadata must + // survive cleanup. Volatile loads keep the control loop from disappearing. + ASSERT_EXIT( + { + auto result = compileAndOptimizeIR(R"( +define i64 @user_unroll_guard(ptr %data, i64 %count) { +entry: + br label %loop + +loop: + %index = phi i64 [ 0, %entry ], [ %next, %loop ] + %total = phi i64 [ 0, %entry ], [ %updated, %loop ] + %address = getelementptr i64, ptr %data, i64 %index + %value = load volatile i64, ptr %address + %updated = add i64 %total, %value + %next = add i64 %index, 1 + %again = icmp ult i64 %next, %count + br i1 %again, label %loop, label %exit, !llvm.loop !0 + +exit: + ret i64 %updated +} + +!0 = distinct !{!0, !1, !2} +!1 = !{!"llvm.loop.unroll.disable"} +!2 = !{!"codon.test.preserve"} +)"); + ASSERT_NE(nullptr, result.module); + auto *function = result.module->getFunction("user_unroll_guard"); + ASSERT_NE(nullptr, function); + llvm::DominatorTree dominators(*function); + llvm::LoopInfo loops(dominators); + ASSERT_FALSE(loops.empty()); + auto *identifier = (*loops.begin())->getLoopID(); + ASSERT_NE(nullptr, identifier); + bool disabled = false; + bool preserved = false; + for (unsigned index = 1; index < identifier->getNumOperands(); ++index) { + auto *node = llvm::dyn_cast(identifier->getOperand(index)); + if (!node || !node->getNumOperands()) + continue; + auto *name = llvm::dyn_cast(node->getOperand(0)); + if (!name) + continue; + disabled |= name->getString() == "llvm.loop.unroll.disable"; + preserved |= name->getString() == "codon.test.preserve"; + } + EXPECT_TRUE(disabled); + EXPECT_TRUE(preserved); + std::_Exit(HasFailure() ? EXIT_FAILURE : EXIT_SUCCESS); + }, + testing::ExitedWithCode(EXIT_SUCCESS), ""); +} + +TEST(LLVMOptimizationTest, FoldsFilteredGeneratorsWithoutPrematureUnrolling) { + // The LLVM fix must fold the reducer independently of CIR producer/consumer fusion. + for (bool disableFusion : {false, true}) { + SCOPED_TRACE(disableFusion); + ASSERT_EXIT( + { + std::vector disabled; + if (disableFusion) + disabled.push_back("core-pythonic-generator-loop-fusion"); + std::string code = R"( +@export +def filtered_start() -> int: + return sum(filter(lambda value: value % 2 == 0, + (index for index in range(25))), 7) + +@export +def filtered_reject() -> int: + return sum(filter(lambda value: value < 0, + (index for index in range(25)))) + +@export +def unfiltered() -> int: + return sum(index for index in range(25)) + +@export +def consumed() -> int: + total = 0 + for value in filter(lambda value: value % 2 == 0, + (index for index in range(25))): + total += value + return total + +@export +def ordinary(data: Ptr[int]) -> int: + total = 0 + for index in range(4): + total += data[index] + return total + +def yielding(data: Ptr[int]): + for value in range(25): + total = 0 + for index in range(4): + total += data[index] + yield total + value + +@export +def escaping(data: Ptr[int]): + return yielding(data) +)"; + for (int count : {0, 1, 2, 24, 25, 26, 64}) + code += "\n@export\ndef filtered_" + std::to_string(count) + + "() -> int:\n" + " return sum(filter(lambda value: value % 2 == 0, " + "(index for index in range(" + + std::to_string(count) + "))))\n"; + auto compiler = compileAndOptimize(code, disabled); + auto *module = compiler->getLLVMVisitor()->getModule(); + EXPECT_FALSE(llvm::verifyModule(*module, &llvm::errs())); + std::string ir; + llvm::raw_string_ostream output(ir); + module->print(output, nullptr); + // The unfused direct consumer exercises metadata left on non-loop branches. + EXPECT_EQ(ir.find("codon.coro.unroll"), std::string::npos); + auto expectConstant = [&](const std::string &name, int64_t expected) { + auto *function = module->getFunction(name); + ASSERT_NE(nullptr, function); + ASSERT_EQ(function->size(), 1u) << name; + ASSERT_EQ(function->front().size(), 1u) << name; + auto *result = + llvm::dyn_cast(function->front().getTerminator()); + ASSERT_NE(nullptr, result); + auto *value = llvm::dyn_cast(result->getReturnValue()); + ASSERT_NE(nullptr, value) << name; + EXPECT_EQ(value->getSExtValue(), expected) << name; + }; + for (int count : {0, 1, 2, 24, 25, 26, 64}) { + int64_t expected = 0; + for (int index = 0; index < count; index += 2) + expected += index; + expectConstant("filtered_" + std::to_string(count), expected); + } + expectConstant("filtered_start", 163); + expectConstant("filtered_reject", 0); + expectConstant("unfiltered", 300); + if (!disableFusion) + expectConstant("consumed", 156); + // Guard against accidentally disabling ordinary unrolling globally or + // throughout an entire coroutine function instead of just yielding loops. + auto *ordinary = module->getFunction("ordinary"); + ASSERT_NE(nullptr, ordinary); + llvm::DominatorTree dominators(*ordinary); + llvm::LoopInfo loops(dominators); + EXPECT_TRUE(loops.empty()); + bool sawResume = false; + for (auto &function : *module) { + if (!function.getName().contains("yielding") || + !function.getName().ends_with(".resume")) + continue; + sawResume = true; + llvm::DominatorTree resumeDominators(function); + llvm::LoopInfo resumeLoops(resumeDominators); + EXPECT_TRUE(resumeLoops.empty()); + } + EXPECT_TRUE(sawResume); + std::_Exit(HasFailure() ? EXIT_FAILURE : EXIT_SUCCESS); + }, + testing::ExitedWithCode(EXIT_SUCCESS), ""); + } +} + +TEST(LLVMOptimizationTest, FusesNestedGeneratorsIntoHandwrittenLoopStructure) { + ASSERT_EXIT( + { + auto compiler = compileAndOptimize(R"( +def fusion_values(data: Ptr[int], count: int): + for index in range(count): + yield data[index] + +def fusion_positive(items): + for value in items: + if value > 0: + yield value + +def fusion_triple(items): + for value in items: + yield value * 3 + +@export +def fusion_handwritten(data: Ptr[int], count: int) -> int: + total = 0 + for index in range(count): + value = data[index] + if value > 0: + total += value * 3 + return total + +@export +def fusion_builtins(data: Ptr[int], count: int) -> int: + total = 0 + for value in map(lambda value: value * 3, + filter(lambda value: value > 0, fusion_values(data, count))): + total += value + return total + +@export +def fusion_nested(data: Ptr[int], count: int) -> int: + total = 0 + for value in fusion_triple(fusion_positive(fusion_values(data, count))): + total += value + return total + +@export +def fusion_expression(data: Ptr[int], count: int) -> int: + total = 0 + for value in (value * 3 for value in fusion_values(data, count) if value > 0): + total += value + return total + +@export +def fusion_deep(data: Ptr[int], count: int) -> int: + total = 0 + for value in fusion_triple(fusion_positive(fusion_positive( + fusion_positive(fusion_values(data, count))))): + total += value + return total + +def fusion_frame_local(value: int): + local = value + yield __ptr__(local) + +@export +def fusion_frame_address(value: int) -> Ptr[int]: + for address in fusion_frame_local(value): + return address + return Ptr[int]() +)"); + auto *module = compiler->getLLVMVisitor()->getModule(); + EXPECT_FALSE(llvm::verifyModule(*module, &llvm::errs())); + auto structure = [](llvm::Function *function) { + std::vector> result; + llvm::DominatorTree dominators(*function); + llvm::LoopInfo loops(dominators); + for (auto &block : *function) { + EXPECT_LE(loops.getLoopDepth(&block), 1u); + std::vector signature; + signature.push_back(loops.getLoopDepth(&block)); + for (auto &instruction : block) { + signature.push_back(instruction.getOpcode()); + if (auto *call = llvm::dyn_cast(&instruction)) { + auto *callee = call->getCalledFunction(); + EXPECT_TRUE(callee && callee->isIntrinsic() && + !callee->getName().starts_with("llvm.coro.")); + } + EXPECT_FALSE(llvm::isa(instruction)); + } + for (auto *successor : llvm::successors(&block)) + signature.push_back( + std::distance(function->begin(), successor->getIterator())); + result.push_back(std::move(signature)); + } + return result; + }; + auto expected = structure(module->getFunction("fusion_handwritten")); + for (auto *name : + {"fusion_builtins", "fusion_nested", "fusion_expression", "fusion_deep"}) { + SCOPED_TRACE(name); + auto *function = module->getFunction(name); + ASSERT_NE(nullptr, function); + EXPECT_EQ(expected, structure(function)); + } + bool retainsFrame = false; + for (auto *value : *compiler->getModule()) { + auto *function = ir::cast(value); + if (!function || function->getUnmangledName() != "fusion_frame_address") + continue; + std::vector pending{function->getBody()}; + while (!pending.empty()) { + auto *node = pending.back(); + pending.pop_back(); + retainsFrame |= ir::isA(node); + auto children = node->getUsedValues(); + pending.insert(pending.end(), children.begin(), children.end()); + } + } + EXPECT_TRUE(retainsFrame); + std::_Exit(HasFailure() ? EXIT_FAILURE : EXIT_SUCCESS); + }, + testing::ExitedWithCode(EXIT_SUCCESS), ""); +} + +TEST(LLVMOptimizationTest, MergesEquivalentGuardedPointerStates) { + for (bool sameInitial : {true, false}) { + for (bool nullEdge : {true, false}) { + SCOPED_TRACE(sameInitial); + SCOPED_TRACE(nullEdge); + std::string code = R"( +declare ptr @advance(ptr) +define i1 @test(ptr %initial, ptr %other, i64 %count) { +entry: + br label %header +header: + %left = phi ptr [ %initial, %entry ], [ %left.next, %latch ] + %right = phi ptr [ INITIAL, %entry ], [ %right.next, %latch ] + %index = phi i64 [ 0, %entry ], [ %next, %latch ] + %finished = icmp eq i64 %index, %count + br i1 %finished, label %exit, label %pull +pull: + %done = icmp PREDICATE ptr %right, null + br i1 %done, label %latch, label %resume +resume: + %value = call ptr @advance(ptr %left) + br label %latch +latch: + %left.next = phi ptr [ %left, %pull ], [ %value, %resume ] + %right.next = phi ptr [ null, %pull ], [ %value, %resume ] + %next = add i64 %index, 1 + br label %header +exit: + %equal = icmp eq ptr %left, %right + ret i1 %equal +} +)"; + code.replace(code.find("INITIAL"), 7, sameInitial ? "%initial" : "%other"); + code.replace(code.find("PREDICATE"), 9, nullEdge ? "eq" : "ne"); + auto optimized = compileAndOptimizeIR(code); + ASSERT_NE(nullptr, optimized.module); + auto *function = optimized.module->getFunction("test"); + ASSERT_NE(nullptr, function); + bool alwaysEqual = true; + bool returns = false; + for (auto &block : *function) { + if (auto *result = llvm::dyn_cast(block.getTerminator())) { + auto *constant = llvm::dyn_cast(result->getReturnValue()); + alwaysEqual &= constant && constant->isOne(); + returns = true; + } + } + EXPECT_TRUE(returns); + EXPECT_EQ(sameInitial && nullEdge, alwaysEqual); + } + } +} + +TEST(LLVMOptimizationTest, ElidesRepeatedGeneratorPulls) { + ASSERT_EXIT( + { + std::string code = R"( +def values(data: Ptr[int], count: int): + for index in range(count): + yield data[index] + +@export +def return_generator(data: Ptr[int], count: int): + return values(data, count) + +@export +def store_generator(data: Ptr[int], count: int, output: Ptr[Generator[int]]): + output[0] = values(data, count) +)"; + for (int pulls : {16, 32, 64}) { + code += "\n@export\ndef pull_" + std::to_string(pulls) + + "(data: Ptr[int], count: int) -> int:\n" + " items = values(data, count)\n total = 0\n"; + for (int index = 0; index < pulls; ++index) + code += " total += next(items, -1)\n"; + code += " return total\n"; + } + auto compiler = compileAndOptimize(code); + auto *module = compiler->getLLVMVisitor()->getModule(); + EXPECT_FALSE(llvm::verifyModule(*module, &llvm::errs())); + for (int pulls : {16, 32, 64}) { + auto name = "pull_" + std::to_string(pulls); + SCOPED_TRACE(name); + auto *function = module->getFunction(name); + ASSERT_NE(nullptr, function); + for (auto &block : *function) { + for (auto &instruction : block) { + EXPECT_FALSE(llvm::isa(instruction)); + if (auto *call = llvm::dyn_cast(&instruction)) { + auto *callee = call->getCalledFunction(); + EXPECT_TRUE(callee && callee->isIntrinsic() && + !callee->getName().starts_with("llvm.coro.")) + << (callee ? callee->getName().str() : "indirect call"); + } + } + } + } + for (const auto *name : {"return_generator", "store_generator"}) { + SCOPED_TRACE(name); + auto *function = module->getFunction(name); + ASSERT_NE(nullptr, function); + bool allocates = false; + for (auto &block : *function) { + for (auto &instruction : block) { + if (auto *call = llvm::dyn_cast(&instruction)) { + auto *callee = call->getCalledFunction(); + allocates |= callee && callee->getName() == "seq_alloc"; + } + } + } + EXPECT_TRUE(allocates); + } + std::_Exit(HasFailure() ? EXIT_FAILURE : EXIT_SUCCESS); + }, + testing::ExitedWithCode(EXIT_SUCCESS), ""); +} + TEST(LLVMOptimizationTest, RequiresKnownNumpyOwnership) { using ir::transform::numpy::hasOwnedResult; using ir::transform::numpy::NumPyExpr; diff --git a/test/cir/transform/enumerate.cpp b/test/cir/transform/enumerate.cpp new file mode 100644 index 000000000..8257d436a --- /dev/null +++ b/test/cir/transform/enumerate.cpp @@ -0,0 +1,189 @@ +#include "test.h" + +#include + +#include "codon/cir/transform/lowering/imperative.h" +#include "codon/cir/transform/pythonic/enumerate.h" +#include "codon/cir/util/irtools.h" +#include "codon/compiler/compiler.h" +#include "codon/compiler/options.h" + +using namespace codon; + +namespace { +struct LoopCounts { + int generators = 0; + int indexed = 0; + int enumerates = 0; + int globalTargets = 0; +}; + +class EnumerateInspector : public ir::util::Operator { +public: + std::unordered_map counts; + bool markSchedules = false; + + void handle(ir::ForFlow *loop) override { + auto name = getParentFunc()->getUnmangledName(); + ++counts[name].generators; + if (loop->getVar()->isGlobal()) + ++counts[name].globalTargets; + if (markSchedules && name == "enum_parallel") + loop->setParallel(); + if (markSchedules && name == "enum_async") + loop->setAsync(); + } + + void handle(ir::ImperativeForFlow *loop) override { + ++counts[getParentFunc()->getUnmangledName()].indexed; + } + + void handle(ir::CallInstr *call) override { + auto *func = ir::util::getFunc(call->getCallee()); + if (func && func->getName().rfind( + ast::getMangledFunc("std.internal.builtin", "enumerate"), 0) == 0) + ++counts[getParentFunc()->getUnmangledName()].enumerates; + } +}; +} // namespace + +TEST(EnumerateOptimizationTest, RemovesGeneratorsAndPreservesFallbacks) { + auto options = Options::getDefault("build/codon_test"); + options->debug = false; + options->native = false; + Compiler compiler(*options); + llvm::cantFail(compiler.parseCode("enumerate_optimization_test.codon", R"( +import numpy as np + +def values(): + yield 1 + +def enum_generator(items: Generator[int]): + for index, value in enumerate(items, 7): + print(index, value) + +class EnumBase(object): + def __iter__(self) -> Generator[int]: + yield 1 + +class EnumDerived(EnumBase): + def __iter__(self) -> Generator[int]: + yield 2 + +def enum_virtual(items: EnumBase): + for index, value in enumerate(items): + print(index, value) + +class EnumDefault: + def __iter__(self, start: int = 3): + yield start + +def enum_default(items: EnumDefault): + for index, value in enumerate(items): + print(index, value) + +global_pair = (-1, -1) + +def observe_global_pair(): + print(global_pair) + +def enum_global(items: List[int]): + global global_pair + for global_pair in enumerate(items): + index = global_pair[0] + value = global_pair[1] + print(index, value) + observe_global_pair() + +def enum_list(items: List[int]): + for index, value in enumerate(items): + print(index, value) + +def enum_array(items: np.ndarray[int, 1]): + for index, value in enumerate(items): + print(index, value) + +def enum_matrix(items: np.ndarray[int, 2]): + for index, row in enumerate(items): + print(index, row) + +def enum_tuple(items: List[int]): + for pair in enumerate(items): + index = pair[0] + value = pair[1] + print(index, value, pair) + +def enum_escaping(items: List[int]): + iterator = enumerate(items) + for index, value in iterator: + print(index, value) + +def enum_parallel(items: List[int]): + for index, value in enumerate(items): + print(index, value) + +def enum_async(items: List[int]): + for index, value in enumerate(items): + print(index, value) + +def enum_shadowed(items: List[int]): + def enumerate(items): + for value in items: + yield (42, value) + for index, value in enumerate(items): + print(index, value) + +enum_generator(values()) +enum_virtual(EnumDerived()) +enum_default(EnumDefault()) +enum_global([1]) +enum_list([1]) +enum_array(np.arange(2)) +enum_matrix(np.arange(4).reshape(2, 2)) +enum_tuple([1]) +enum_escaping([1]) +enum_parallel([1]) +enum_async([1]) +enum_shadowed([1]) +)")); + + EnumerateInspector before; + before.markSchedules = true; + before.process(compiler.getModule()); + EXPECT_EQ(1, before.counts["enum_generator"].enumerates); + EXPECT_EQ(1, before.counts["enum_list"].enumerates); + EXPECT_EQ(1, before.counts["enum_array"].enumerates); + EXPECT_EQ(1, before.counts["enum_matrix"].enumerates); + EXPECT_EQ(1, before.counts["enum_global"].globalTargets); + EXPECT_EQ(1, before.counts["enum_virtual"].enumerates); + EXPECT_EQ(1, before.counts["enum_default"].enumerates); + + ir::transform::pythonic::EnumerateOptimization enumerate; + enumerate.run(compiler.getModule()); + ir::transform::lowering::ImperativeForFlowLowering lowering; + lowering.run(compiler.getModule()); + + EnumerateInspector after; + after.process(compiler.getModule()); + for (const auto &name : {"enum_generator", "enum_virtual", "enum_default"}) { + SCOPED_TRACE(name); + EXPECT_EQ(0, after.counts[name].enumerates); + EXPECT_EQ(1, after.counts[name].generators); + } + EXPECT_EQ(1, after.counts["enum_global"].globalTargets); + for (const auto &name : {"enum_list", "enum_array", "enum_matrix"}) { + SCOPED_TRACE(name); + EXPECT_EQ(0, after.counts[name].enumerates); + EXPECT_EQ(0, after.counts[name].generators); + EXPECT_EQ(1, after.counts[name].indexed); + } + for (const auto &name : + {"enum_tuple", "enum_escaping", "enum_parallel", "enum_async", "enum_global"}) { + SCOPED_TRACE(name); + EXPECT_EQ(1, after.counts[name].enumerates); + EXPECT_EQ(1, after.counts[name].generators); + EXPECT_EQ(0, after.counts[name].indexed); + } + EXPECT_EQ(0, after.counts["enum_shadowed"].enumerates); + EXPECT_EQ(1, after.counts["enum_shadowed"].generators); +} diff --git a/test/core/bltin.codon b/test/core/bltin.codon index b86a802be..22b4407cf 100644 --- a/test/core/bltin.codon +++ b/test/core/bltin.codon @@ -902,13 +902,17 @@ def test_num_from_str(): assert float("é3.14é"[1:-1]) == 3.14 assert float("Ω3.14Ω"[1:-1]) == 3.14 assert float("😀3.14😀"[1:-1]) == 3.14 + assert float("3.14trailing"[:4]) == 3.14 + assert float("x3.14trailing"[1:5]) == 3.14 assert complex(" (1+2j) ") == 1+2j assert complex("é1+2jé"[1:-1]) == 1+2j assert complex("Ω1+2jΩ"[1:-1]) == 1+2j assert complex("😀1+2j😀"[1:-1]) == 1+2j + assert complex("1+2jtrailing"[:4]) == 1+2j + assert complex("x1+2jtrailing"[1:5]) == 1+2j - for value in ("", "é"): + for value in ("", "é", "1\x00", "1e", "1e+", "+"): try: float(value) assert False @@ -921,6 +925,47 @@ def test_num_from_str(): except ValueError: pass + for value in ("1+2j\x00", "1+", "(1+2j", "1+2j)"): + try: + complex(value) + assert False + except ValueError: + pass + +@test +def test_int_from_wide_storage(): + for marker in ("x", "é", "Ω", "😀"): + for text, base, expected in ( + (" -123\t", 10, -123), + ("0x_FF", 0, 255), + ("0X_FF", 16, 255), + ("0b_101", 2, 5), + ("0o_17", 8, 15), + ("z", 36, 35), + ("9223372036854775807", 10, 9223372036854775807), + ("-9223372036854775808", 10, -9223372036854775808), + ): + value = (marker + text + marker)[1:-1] + assert int(value, base) == expected + + for text, base in (("", 10), ("_1", 10), ("1_", 10), + ("1__2", 10), ("1\x00", 10), ("0x", 0), + ("0x__1", 16), ("01", 0)): + value = (marker + text + marker)[1:-1] + try: + int(value, base) + assert False + except ValueError: + pass + + for text in ("9223372036854775808", "-9223372036854775809"): + value = (marker + text + marker)[1:-1] + try: + int(value) + assert False + except OverflowError: + pass + @test def test_num_from_str_extended(): # Basic decimal @@ -1782,6 +1827,7 @@ test_reversed() test_divmod() test_pow() test_num_from_str() +test_int_from_wide_storage() test_num_from_str_extended() test_files(open) import gzip diff --git a/test/core/containers.codon b/test/core/containers.codon index c3f3fabe1..6a0fd8161 100644 --- a/test/core/containers.codon +++ b/test/core/containers.codon @@ -221,6 +221,43 @@ def test_list(): assert [a for a in l2] == [1, 2, 1, 2] assert 2 * [1, 2] == l2 + for count in (-2, 0, 1, 2, 17, 129): + assert [7] * count == [7 for index in range(max(0, count))] + assert [1.25] * count == [1.25 for index in range(max(0, count))] + assert ["value"] * count == ["value" for index in range(max(0, count))] + assert [(1, "x")] * count == [(1, "x") for index in range(max(0, count))] + assert [None] * count == [None for index in range(max(0, count))] + assert List[int]() * count == [] + assert [1, 2, 3] * count == [value for index in range(max(0, count)) for value in (1, 2, 3)] + assert [1.25, 2.5] * count == [value for index in range(max(0, count)) for value in (1.25, 2.5)] + assert ["first", "second"] * count == [value for index in range(max(0, count)) for value in ("first", "second")] + assert [(1, "x"), (2, "y")] * count == [value for index in range(max(0, count)) for value in ((1, "x"), (2, "y"))] + assert [None, None] * count == [None for index in range(max(0, 2 * count))] + assert List[int]() * 0x7FFFFFFFFFFFFFFF == [] + source = [1] + repeated = [source] * 7 + source.append(2) + assert all(value is source and value == [1, 2] for value in repeated) + repeated.append([3]) + assert len(repeated) == 8 + assert source == [1, 2] + + other = [3] + original = [source, other] + repeated = 17 * original + assert repeated is not original + assert len(repeated) == 34 + for index in range(17): + assert repeated[2 * index] is source + assert repeated[2 * index + 1] is other + repeated[0] = [99] + other.append(4) + repeated.append([5]) + assert original == [[1, 2], [3, 4]] + assert repeated[0] == [99] and repeated[2] is source + assert repeated[1] == [3, 4] and repeated[-1] == [5] + assert len(repeated) == 35 + l1 = [i*2 for i in range(3)] l1.insert(0, 99) l1[0] += 1 diff --git a/test/core/generators.codon b/test/core/generators.codon index b46ce1adf..1375ebed1 100644 --- a/test/core/generators.codon +++ b/test/core/generators.codon @@ -87,6 +87,70 @@ for i in range(10): print(fadder.send(-1.0)) # EXPECT: 45.0 +@test +def test_generator_send(): + def exchange(events): + value = (yield) + events.append(value) + yield value + 1 + events.append(99) + + events = list[int]() + items = exchange(events) + next(items) + assert items.send(10) == 11 + for repeat in range(3): + try: + items.send(100) + assert False + except StopIteration: + pass + assert events == [10, 99] + assert next(items, -1) == -1 + + def source(count): + for value in range(count): + yield value + + for count in [0, 1]: + items = source(count) + assert list(items) == list(range(count)) + for repeat in range(3): + try: + items.send(100) + assert False + except StopIteration: + pass + + items = source(0) + try: + items.send(100) + assert False + except StopIteration: + pass + + def empty_values(events): + events.append(1) + yield None + events.append(2) + yield None + events.append(3) + + events = list[int]() + items = empty_values(events) + next(items) + items.send(None) + assert events == [1, 2] + for repeat in range(3): + try: + items.send(None) + assert False + except StopIteration: + pass + assert events == [1, 2, 3] +test_generator_send() + + @test def test_generator_in_finally(): def foo(n): @@ -108,3 +172,315 @@ def test_generator_in_finally(): assert e.message == 'not n' assert b test_generator_in_finally() + +@test +def test_generator_loop_fusion(): + def source(data): + for value in data: + yield value + + def positive(items): + for value in items: + if value > 0: + yield value + + def triple(items): + for value in items: + yield value * 3 + + def builtin_total(data): + total = 0 + for value in map(lambda value: value * 3, + filter(lambda value: value > 0, source(data))): + total += value + return total + + def nested_total(data): + total = 0 + for value in triple(positive(source(data))): + total += value + return total + + def expression_total(data): + total = 0 + for value in (value * 3 for value in source(data) if value > 0): + total += value + return total + + for data in [[], [0], [1], [-1], [1, 2, 3], [-1, -2, -3], + [-3, 0, 4, -2, 5, 1]]: + expected = 0 + for value in data: + if value > 0: + expected += value * 3 + assert builtin_total(data) == expected + assert nested_total(data) == expected + assert expression_total(data) == expected + + items = source([1, 2, 3]) + totals = [] + for repeat in range(3): + total = 0 + for value in items: + total += value + totals.append(total) + assert totals == [6, 0, 0] + + def first(data: list[int]): + for value in positive(source(data)): + return value + return -1 + assert first([-1, 2, 3]) == 2 + assert first([]) == -1 + + total = 0 + for value in source([1, 2, 3, 4]): + if value == 2: + continue + if value == 4: + break + total += value + assert total == 4 +test_generator_loop_fusion() + +@test +def test_generator_fusion_order_and_exceptions(): + def capture(events, label, value): + events.append(label) + return value + + def source(events, count): + for value in range(count): + events.append(10 + value) + yield value + events.append(20 + value) + + def selected(events, items): + for value in items: + events.append(30 + value) + if value % 2 == 0: + yield value + + events = list[int]() + count = 3 + items = selected(capture(events, 1, events), + source(events, capture(events, 2, count))) + count = 0 + events.append(3) + for value in items: + events.append(40 + value) + assert events == [1, 2, 3, 10, 30, 40, 20, 11, 31, 21, 12, 32, 42, 22] + + def raising(value): + if value == 2: + raise ValueError('callback') + return value * 3 + + events = list[int]() + try: + for value in map(raising, source(events, 4)): + events.append(40 + value) + assert False + except ValueError as error: + assert error.message == 'callback' + assert events == [10, 40, 20, 11, 43, 21, 12] + + def protected(events): + try: + for value in range(3): + yield value + except ValueError: + events.append(99) + + events = list[int]() + try: + for value in protected(events): + raise ValueError('consumer') + assert False + except ValueError as error: + assert error.message == 'consumer' + assert events == [] + + def truncated(): + for value in range(5): + if value == 3: + return + yield value + + total = 0 + for value in truncated(): + total += value + assert total == 3 +test_generator_fusion_order_and_exceptions() + +@test +def test_generator_exhaustion(): + def values(): + yield 1 + yield 2 + + items = values() + assert not hasattr(items, 'done') + assert iter(items).__raw__() == items.__raw__() + assert items.__next__() == 1 + assert next(items) == 2 + for attempt in range(3): + assert next(items, -1) == -1 + assert list(items) == [] + try: + items.__next__() + assert False + except StopIteration: + pass + + items = values() + assert list(items) == [1, 2] + assert list(items) == [] + assert next(items, 0) == 0 + assert min(items, default=0) == 0 + assert max(items, default=0) == 0 + + items = values() + seen = [] + for value in items: + seen.append(value) + assert list(items) == [2] + assert seen == [1] + + items = values() + seen = [] + for value in items: + seen.append(value) + assert list(items) == [2] + continue + assert seen == [1] + +test_generator_exhaustion() + +@test +def test_zip_generators(): + def values(count): + for value in range(count): + yield value + + items = [1, 2, 3] + pairs = zip(items, items) + assert list(pairs) == [(1, 1), (2, 2), (3, 3)] + assert list(zip(pairs, pairs)) == [] + assert list(zip()) == [] + assert list(zip([1, 2])) == [(1,), (2,)] + assert list(zip(['a', 'b'], [1, 2], [True, False])) == [('a', 1, True), ('b', 2, False)] + + shared = values(5) + assert list(zip(shared, shared)) == [(0, 1), (2, 3)] + assert next(shared, -1) == -1 + shared = values(8) + assert list(zip(shared, shared, shared)) == [(0, 1, 2), (3, 4, 5)] + + first = values(1) + second = values(3) + assert list(zip(first, second)) == [(0, 0)] + assert next(second) == 1 + assert list(second) == [2] + + first = values(0) + second = values(3) + assert list(zip(first, second)) == [] + assert next(second) == 0 + + first = values(3) + second = values(1) + third = values(3) + assert list(zip(first, second, third)) == [(0, 0, 0)] + assert next(first) == 2 + assert next(third) == 1 + + nullable: List[Optional[int]] = [None, 1] + rows = list(zip(nullable, ['a', 'b'])) + assert len(rows) == 2 + assert rows[0][0] is None and rows[0][1] == 'a' + assert rows[1][0] == 1 and rows[1][1] == 'b' + + class Payload: + data: List[int] + + def __init__(self, data: List[int]): + self.data = data + + payload = Payload([10]) + rows = list(zip([payload], ['value'], [(1, NoneType())])) + assert rows[0][0] is payload + rows[0][0].data.append(20) + assert payload.data == [10, 20] + assert rows[0][1] == 'value' and rows[0][2][0] == 1 + +test_zip_generators() + +@test +def test_zip_advancement_order(): + def values(events, tag, count): + events.append(tag + '-start') + for value in range(count): + events.append(tag + str(value)) + yield value + events.append(tag + '-end') + + events = [] + first = values(events, 'a', 1) + second = values(events, 'b', 3) + pairs = zip(first, second) + assert events == [] + assert next(pairs) == (0, 0) + assert events == ['a-start', 'a0', 'b-start', 'b0'] + assert list(pairs) == [] + assert events == ['a-start', 'a0', 'b-start', 'b0', 'a-end'] + assert next(second) == 1 + assert events[-1] == 'b1' + + def raising(): + yield 1 + raise ValueError('iterator failed') + + events = [] + pairs = zip(raising(), values(events, 'b', 3)) + assert next(pairs) == (1, 0) + try: + next(pairs) + assert False + except ValueError as error: + assert error.message == 'iterator failed' + assert events == ['b-start', 'b0'] + + def void_values(): + yield + yield + + assert len(list(zip(void_values(), range(3)))) == 2 + +test_zip_advancement_order() + +@test +def test_exhausted_generator_pipeline(): + def values(): + yield 1 + yield 2 + + def collect(value, result): + result.append(value) + + items = values() + result = [] + items |> collect(..., result) + items |> collect(..., result) + assert result == [1, 2] + + def drain(value, items, result): + result.append(value) + assert list(items) == [2] + + items = values() + result = [] + items |> drain(..., items, result) + assert result == [1] + +test_exhausted_generator_pipeline() diff --git a/test/core/numerics.codon b/test/core/numerics.codon index aed391b6b..50a2650b4 100644 --- a/test/core/numerics.codon +++ b/test/core/numerics.codon @@ -8,6 +8,22 @@ NINF = -math.inf def close(a: float, b: float, epsilon: float = 1e-7): return abs(a - b) <= epsilon +@test +def test_factorial_table(): + expected = 1 + for value in range(21): + if value: + expected *= value + assert math.factorial(value) == expected + for value in (-9223372036854775807 - 1, -1, 21, 9223372036854775807): + try: + math.factorial(value) + assert False + except ValueError as error: + assert str(error) == "factorial is only supported for 0 <= x <= 20" + +test_factorial_table() + if __py_numerics__: ### Python semantics tests diff --git a/test/main.cpp b/test/main.cpp index 24df66fd0..fd072ff6a 100644 --- a/test/main.cpp +++ b/test/main.cpp @@ -536,6 +536,7 @@ INSTANTIATE_TEST_SUITE_P( testing::Values( "transform/canonical.codon", "transform/dict_opt.codon", + "transform/enumerate.codon", "transform/escapes.codon", "transform/folding.codon", "transform/for_lowering.codon", diff --git a/test/numpy/test_indexing.codon b/test/numpy/test_indexing.codon index 2b954e53d..896f43a20 100644 --- a/test/numpy/test_indexing.codon +++ b/test/numpy/test_indexing.codon @@ -1,6 +1,54 @@ import numpy as np from numpy import * +@test +def test_enumerate_indexing(): + data = np.arange(12).reshape(3, 4) + for index, row in enumerate(data, 5): + assert np.array_equal(row, data[index - 5]) + row[0] = 100 + index + assert list(data[:, 0]) == [105, 106, 107] + + for index, column in enumerate(data.T): + assert np.array_equal(column, data[:, index]) + + result = [] + for index, value in enumerate(data[1, ::-2], -1): + result.append((index, value)) + index = 99 + continue + assert result == [(-1, 7), (0, 5)] + + for index, value in enumerate(np.arange(0)): + assert False + for index, row in enumerate(np.empty((0, 3))): + assert False + for index, row in enumerate(np.empty((2, 0)), -3): + assert row.shape == (0,) + assert index == -3 or index == -2 + + def source(events): + events.append('source') + return np.arange(3) + + def start(events): + events.append('start') + return 8 + + events = [] + result = [] + for index, value in enumerate(source(events), start(events)): + result.append((index, value)) + if value == 1: + break + else: + assert False + assert events == ['source', 'start'] + assert result == [(8, 0), (9, 1)] + assert index == 9 and value == 1 + +test_enumerate_indexing() + @test def test_basic_indexing(): #Single element indexing diff --git a/test/parser/llvm.codon b/test/parser/llvm.codon index 46ab55cc1..106f94024 100644 --- a/test/parser/llvm.codon +++ b/test/parser/llvm.codon @@ -179,13 +179,15 @@ def foo(): yield 3 z = foo() y = z.__iter__() +print hasattr(y, 'destroy') #: False +print hasattr(y, '_destroy') #: True print str(y.__raw__())[:2] #: 0x print y.__done__() #: False print y.__promise__()[0] #: 0 y.__resume__() print y.__repr__()[:16] #: 0: + expected += separator + expected += piece + assert separator.join(pieces) == expected + @test def test_repr(): diff --git a/test/transform/enumerate.codon b/test/transform/enumerate.codon new file mode 100644 index 000000000..144e09455 --- /dev/null +++ b/test/transform/enumerate.codon @@ -0,0 +1,224 @@ +def enumerate_values(count: int): + for value in range(count): + yield value * 3 + +@test +def test_enumerate_control_flow(): + result = [] + for index, value in enumerate(enumerate_values(5), -2): + result.append((index, value)) + index = 99 + if value == 3: + continue + if value == 9: + break + else: + assert False + assert result == [(-2, 0), (-1, 3), (0, 6), (1, 9)] + + result = [] + for index, value in enumerate([3, 6, 9], 4): + try: + if value == 6: + continue + result.append((index, value)) + finally: + result.append((index, -value)) + else: + result.append((-1, -1)) + assert result == [(4, 3), (4, -3), (5, -6), (6, 9), (6, -9), (-1, -1)] + assert index == 6 and value == 9 + + for index, value in enumerate(enumerate_values(0), 100): + assert False + assert index == 6 and value == 9 + + def first_value(): + for index, value in enumerate(enumerate_values(3), 8): + return index + value + return -1 + assert first_value() == 8 + + result = [] + for index, value in enumerate(range(3), 9223372036854775807): + result.append(index) + assert result == [9223372036854775807, -9223372036854775808, -9223372036854775807] + +test_enumerate_control_flow() + +class EnumerateIterable: + events: List[str] + + def __init__(self, events): + self.events = events + + def __iter__(self, count: int = 2): + self.events.append('iter') + return enumerate_values(count) + +@test +def test_enumerate_evaluation(): + def source(events): + events.append('source') + return EnumerateIterable(events) + + def start(events): + events.append('start') + return -7 + + events = [] + result = [] + for index, value in enumerate(source(events), start(events)): + result.append((index, value)) + assert events == ['source', 'start', 'iter'] + assert result == [(-7, 0), (-6, 3)] + + def fail(events): + events.append('start') + raise ValueError('start') + return 0 + + events = [] + try: + for index, value in enumerate(source(events), fail(events)): + assert False + assert False + except ValueError: + pass + assert events == ['source', 'start'] + + def grow(values): + values.append(30) + return 5 + + values = [10, 20] + result = [] + for index, value in enumerate(values, grow(values)): + result.append((index, value)) + values.append(99) + assert result == [(5, 10), (6, 20), (7, 30)] + +test_enumerate_evaluation() + +@test +def test_enumerate_unpacking(): + result = [] + for index, (left, right) in enumerate([(2, 3), (4, 5)]): + result.append((index, left + right)) + assert result == [(0, 5), (1, 9)] + + result = [] + for outer, left in enumerate([10, 20]): + for inner, right in enumerate(enumerate_values(2), 3): + result.append((outer, inner, left + right)) + assert result == [(0, 3, 10), (0, 4, 13), (1, 3, 20), (1, 4, 23)] + + result = [] + for outer, (inner, value) in enumerate(enumerate([4, 5], 8), 2): + result.append((outer, inner, value)) + assert result == [(2, 8, 4), (3, 9, 5)] + + result = [] + for index, value in enumerate('abc', 2): + result.append((index, value)) + assert result == [(2, 'a'), (3, 'b'), (4, 'c')] + + result = [] + for value, value in enumerate([4, 5]): + result.append(value) + assert result == [4, 5] + +test_enumerate_unpacking() + +@test +def test_enumerate_fallbacks(): + pairs = [] + for pair in enumerate([4, 5], 7): + index = pair[0] + value = pair[1] + pairs.append(pair) + assert index + value == pair[0] + pair[1] + assert pairs == [(7, 4), (8, 5)] + assert pair == (8, 5) + + iterator = enumerate(enumerate_values(2), 3) + result = [] + for index, value in iterator: + result.append((index, value)) + assert result == [(3, 0), (4, 3)] + +test_enumerate_fallbacks() + +@test +def test_enumerate_shadowing(): + def enumerate(values): + return [(99, value) for value in values] + + result = [] + for index, value in enumerate([4, 5]): + result.append((index, value)) + assert result == [(99, 4), (99, 5)] + +test_enumerate_shadowing() + +enumerate_global_pair = (-1, -1) + +@noinline +def observe_enumerate_global_pair(): + return enumerate_global_pair + +@test +def test_enumerate_global_target(): + global enumerate_global_pair + observed = [] + for enumerate_global_pair in enumerate([10, 20], 3): + index = enumerate_global_pair[0] + value = enumerate_global_pair[1] + observed.append(observe_enumerate_global_pair()) + assert observe_enumerate_global_pair() == (index, value) + assert observed == [(3, 10), (4, 20)] + assert observe_enumerate_global_pair() == (4, 20) + +test_enumerate_global_target() + +class EnumerateBase(object): + events: List[str] + + def __init__(self, events: List[str]): + self.events = events + + def __iter__(self) -> Generator[int]: + self.events.append('base') + yield 10 + +class EnumerateDerived(EnumerateBase): + def __iter__(self) -> Generator[int]: + self.events.append('derived') + yield 20 + yield 30 + +@noinline +def enumerate_virtual_values(items: EnumerateBase, start: int): + result = [] + for index, value in enumerate(items, start): + result.append((index, value)) + return result + +@test +def test_enumerate_virtual_iterator(): + def source(events): + events.append('source') + return EnumerateDerived(events) + + def start(events): + events.append('start') + return -3 + + events = [] + assert enumerate_virtual_values(EnumerateBase(events), 2) == [(2, 10)] + assert events == ['base'] + events.clear() + assert enumerate_virtual_values(source(events), start(events)) == [(-3, 20), (-2, 30)] + assert events == ['source', 'start', 'derived'] + +test_enumerate_virtual_iterator() From f9f25347c1b56937101d1e2b44c609240c38ada8 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sat, 26 Sep 2026 17:11:43 -0400 Subject: [PATCH 02/12] Update sort algorithm selection --- codon/runtime/numpy/sort.cpp | 60 +++++++++++++++---------- stdlib/internal/sort.codon | 50 ++++++++++++++++++--- stdlib/numpy/sorting.codon | 32 +------------- test/core/containers.codon | 44 ++++++++++++++++++ test/numpy/test_sorting.codon | 26 +++++++++++ test/stdlib/itertools_test.codon | 60 +++++++++++++++++++++++++ test/stdlib/sort_test.codon | 76 ++++++++++++++++++++++++++++++++ 7 files changed, 288 insertions(+), 60 deletions(-) diff --git a/codon/runtime/numpy/sort.cpp b/codon/runtime/numpy/sort.cpp index 633810aa6..17ea5713c 100644 --- a/codon/runtime/numpy/sort.cpp +++ b/codon/runtime/numpy/sort.cpp @@ -3,38 +3,50 @@ #include "codon/runtime/lib.h" #include "hwy/contrib/sort/vqsort-inl.h" -SEQ_FUNC void cnp_sort_int16(int16_t *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); -} +#include +#include -SEQ_FUNC void cnp_sort_uint16(uint16_t *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); +namespace { +template void sortPrimitive(T *data, int64_t size) { +#ifdef __APPLE__ + // macOS ARM64 benchmarks with 100M random/duplicate-heavy elements favored + // libc++ for these 64-bit types and Highway for 16/32-bit types. + if constexpr (std::is_same_v || std::is_same_v || + std::is_same_v) { + if constexpr (std::is_same_v) { + // NaNs violate std::sort's default ordering; Highway places them last. + // This O(n) scan retains libc++'s specialized default-comparator path. + // Signed zeros compare equivalent and do not require the fallback. + // libc++ sort is significantly faster even with the initial scan. + if (std::any_of(data, data + size, [](T value) { return value != value; })) { + hwy::VQSort(data, size, hwy::SortAscending()); + return; + } + } + std::sort(data, data + size); + return; + } +#endif + hwy::VQSort(data, size, hwy::SortAscending()); } +} // namespace -SEQ_FUNC void cnp_sort_int32(int32_t *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); -} +SEQ_FUNC void cnp_sort_int16(int16_t *data, int64_t n) { sortPrimitive(data, n); } -SEQ_FUNC void cnp_sort_uint32(uint32_t *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); -} +SEQ_FUNC void cnp_sort_uint16(uint16_t *data, int64_t n) { sortPrimitive(data, n); } -SEQ_FUNC void cnp_sort_int64(int64_t *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); -} +SEQ_FUNC void cnp_sort_int32(int32_t *data, int64_t n) { sortPrimitive(data, n); } -SEQ_FUNC void cnp_sort_uint64(uint64_t *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); -} +SEQ_FUNC void cnp_sort_uint32(uint32_t *data, int64_t n) { sortPrimitive(data, n); } -SEQ_FUNC void cnp_sort_uint128(hwy::uint128_t *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); -} +SEQ_FUNC void cnp_sort_int64(int64_t *data, int64_t n) { sortPrimitive(data, n); } -SEQ_FUNC void cnp_sort_float32(float *data, int64_t n) { - hwy::VQSort(data, n, hwy::SortAscending()); -} +SEQ_FUNC void cnp_sort_uint64(uint64_t *data, int64_t n) { sortPrimitive(data, n); } + +SEQ_FUNC void cnp_sort_float32(float *data, int64_t n) { sortPrimitive(data, n); } -SEQ_FUNC void cnp_sort_float64(double *data, int64_t n) { +SEQ_FUNC void cnp_sort_float64(double *data, int64_t n) { sortPrimitive(data, n); } + +SEQ_FUNC void cnp_sort_uint128(hwy::uint128_t *data, int64_t n) { hwy::VQSort(data, n, hwy::SortAscending()); } diff --git a/stdlib/internal/sort.codon b/stdlib/internal/sort.codon index c0bcf5d7f..c996274a2 100644 --- a/stdlib/internal/sort.codon +++ b/stdlib/internal/sort.codon @@ -6,6 +6,47 @@ from algorithms.heapsort import heap_sort_inplace from algorithms.qsort import qsort_inplace from algorithms.timsort import tim_sort_inplace +from C import cnp_sort_int16(cobj, int) +from C import cnp_sort_uint16(cobj, int) +from C import cnp_sort_int32(cobj, int) +from C import cnp_sort_uint32(cobj, int) +from C import cnp_sort_int64(cobj, int) +from C import cnp_sort_uint64(cobj, int) +from C import cnp_sort_uint128(cobj, int) +from C import cnp_sort_float32(cobj, int) +from C import cnp_sort_float64(cobj, int) + +def _try_native_sort(start: Ptr[T], n: int, T: type) -> bool: + if T is int or T is i64: + cnp_sort_int64(start.as_byte(), n) + elif T is i16: + cnp_sort_int16(start.as_byte(), n) + elif T is u16: + cnp_sort_uint16(start.as_byte(), n) + elif T is i32: + cnp_sort_int32(start.as_byte(), n) + elif T is u32: + cnp_sort_uint32(start.as_byte(), n) + elif T is u64: + cnp_sort_uint64(start.as_byte(), n) + elif T is u128: + cnp_sort_uint128(start.as_byte(), n) + elif T is float32: + cnp_sort_float32(start.as_byte(), n) + elif T is float: + cnp_sort_float64(start.as_byte(), n) + else: + return False + return True + +def _try_native_sort_list(start: Ptr[T], n: int, T: type) -> bool: + if T is float or T is float32: + for index in range(n): + value = start[index] + if value != value or value == 0: + return False + return _try_native_sort(start, n) + def sorted( v: Generator[T], key=Optional[int](), @@ -65,13 +106,10 @@ class List: ): if isinstance(key, Optional): if algorithm == "auto": - # Python uses Timsort in all cases, but if we - # know stability does not matter (i.e. sorting - # primitive type with no key), we will use - # faster PDQ instead. PDQ is ~50% faster than - # Timsort for sorting 1B 64-bit ints. if self: - if _is_pdq_compatible(self[0]): + if _try_native_sort_list(self._ptr, self._len): + pass + elif _is_pdq_compatible(self[0]): pdq_sort_inplace(self, lambda x: x) else: tim_sort_inplace(self, lambda x: x) diff --git a/stdlib/numpy/sorting.codon b/stdlib/numpy/sorting.codon index c48a50b31..6b8c3bef6 100644 --- a/stdlib/numpy/sorting.codon +++ b/stdlib/numpy/sorting.codon @@ -2,6 +2,7 @@ from .ndarray import ndarray from .routines import array, asarray, empty, zeros +from internal.sort import _try_native_sort import util import internal.static as static @@ -487,37 +488,8 @@ def _pdq_sort( begin = pivot_pos + 1 leftmost = False -# C stubs for vectorized quicksort from Highway -from C import cnp_sort_int16(cobj, int) -from C import cnp_sort_uint16(cobj, int) -from C import cnp_sort_int32(cobj, int) -from C import cnp_sort_uint32(cobj, int) -from C import cnp_sort_int64(cobj, int) -from C import cnp_sort_uint64(cobj, int) -from C import cnp_sort_uint128(cobj, int) -from C import cnp_sort_float32(cobj, int) -from C import cnp_sort_float64(cobj, int) - def quicksort(start: Ptr[T], n: int, T: type): - if T is int: - cnp_sort_int64(start.as_byte(), n) - elif T is i16: - cnp_sort_int16(start.as_byte(), n) - elif T is u16: - cnp_sort_uint16(start.as_byte(), n) - elif T is i32: - cnp_sort_int32(start.as_byte(), n) - elif T is u32: - cnp_sort_uint32(start.as_byte(), n) - elif T is i64: - cnp_sort_int64(start.as_byte(), n) - elif T is u64: - cnp_sort_uint64(start.as_byte(), n) - elif T is float32: - cnp_sort_float32(start.as_byte(), n) - elif T is float: - cnp_sort_float64(start.as_byte(), n) - else: + if not _try_native_sort(start, n): _pdq_sort(start, 0, n, _floor_log2(n), True) def _apartial_insertion_sort(arr: Ptr[T], tosort: Ptr[int], begin: int, end: int, T: type): diff --git a/test/core/containers.codon b/test/core/containers.codon index 6a0fd8161..0cb02b971 100644 --- a/test/core/containers.codon +++ b/test/core/containers.codon @@ -478,6 +478,50 @@ def test_list(): assert l9 == ["a", "b", "c"] test_list() +@test +def test_list_index_errors(): + for size in (0, 1, 3): + values = list(range(size)) + for index in range(-size, size): + assert values[index] == index % size + for index in (-9223372036854775807 - 1, -size - 1, size, 9223372036854775807): + try: + values[index] + assert False, "out-of-range read succeeded" + except IndexError as error: + assert error.message == "list index out of range" + try: + values[index] = 99 + assert False, "out-of-range assignment succeeded" + except IndexError as error: + assert error.message == "list assignment index out of range" + try: + del values[index] + assert False, "out-of-range deletion succeeded" + except IndexError as error: + assert error.message == "list assignment index out of range" + assert values == list(range(size)) + + def paired_values(data): + for index in range(len(data)): + yield data[index] + yield data[index] * 3 + + assert list(paired_values(list[int]())) == [] + assert list(paired_values([2, 5])) == [2, 6, 5, 15] + values = [2, 5] + items = paired_values(values) + assert next(items) == 2 + values[0] = 7 + assert next(items) == 21 + values.clear() + try: + next(items) + assert False, "generator missed list mutation" + except IndexError as error: + assert error.message == "list index out of range" +test_list_index_errors() + @test def test_setslice(): l = [0, 1] diff --git a/test/numpy/test_sorting.codon b/test/numpy/test_sorting.codon index 8f87227a6..062ce92cd 100644 --- a/test/numpy/test_sorting.codon +++ b/test/numpy/test_sorting.codon @@ -63,6 +63,28 @@ def test_sorts(dtype: type): for length in (0, 1, 10, 101, 1111, 12345, 1000000): test_sort(g, length, kind, dtype) +@test +def test_sort_float_specials(dtype: type): + from math import isnan + nan = float("nan") + inf = float("inf") + for values, expected in (([3.0, nan, -2.0, inf, -inf, nan], + [-inf, -2.0, 3.0, inf, nan, nan]), + ([3.0, -0.0, 0.0, -2.0], [-2.0, 0.0, 0.0, 3.0]), + ([inf, 3.0, -inf, -2.0], [-inf, -2.0, 3.0, inf]), + ([nan], [nan]), ([-0.0], [0.0])): + data = np.array(values, dtype=dtype) + inplace = data.copy() + inplace.sort(kind="quick") + for actual in (np.sort(data, kind="quick"), inplace): + assert len(actual) == len(expected) + for index, reference in enumerate(expected): + value = float(actual[index]) + assert value == reference or (isnan(value) and isnan(reference)) + for index, reference in enumerate(values): + value = float(data[index]) + assert value == reference or (isnan(value) and isnan(reference)) + def check_partitioned(vec, kth): if isinstance(kth, int): return check_partitioned(vec, (kth, )) @@ -168,12 +190,16 @@ test_sorts(int) test_sorts(u64) test_sorts(u32) test_sorts(i32) +test_sorts(u16) +test_sorts(i16) test_sorts(u8) test_sorts(i8) test_sorts(u128) test_sorts(i128) test_sorts(float) test_sorts(float32) +test_sort_float_specials(float) +test_sort_float_specials(float32) test_sorts(complex) test_sorts(complex64) test_partition() diff --git a/test/stdlib/itertools_test.codon b/test/stdlib/itertools_test.codon index 6524434ed..3d5270251 100644 --- a/test/stdlib/itertools_test.codon +++ b/test/stdlib/itertools_test.codon @@ -1207,6 +1207,66 @@ def test_islice_from_cpython(): events.append(index) yield index + def selected(value, period, predicates): + predicates.append(value) + return (value + 1) % period == 0 + + for count in (0, 1, 2, 7, 8, 9, 15, 16, 17, 31, 32, 33, 129): + for period in (1, 2, 7, 300): + matches = [index for index in range(count) if (index + 1) % period == 0] + for stop in (0, 1, 2, 7, 8, 9, 16, 17, count, count + 3): + expected = matches[:stop] + consumed = (0 if stop == 0 else matches[stop - 1] + 1 + if stop <= len(matches) else count) + for explicit_loop in (False, True): + events = list[int]() + predicates = list[int]() + source = observed(events, count) + filtered = (value for value in source + if selected(value, period, predicates)) + sliced = islice(filtered, stop) + assert events == [] and predicates == [] + if explicit_loop: + total = 0 + for value in sliced: + total += value + else: + total = sum(sliced) + assert total == sum(expected) + assert events == list(range(consumed)) + assert predicates == events + assert list(sliced) == [] + following = next(source, -1) + assert following == (consumed if consumed < count else -1) + + def raises_after_prefix(limit): + for value in range(1, limit + 1): + yield value + raise ValueError('past prefix') + + def checked_predicate(value, limit): + if value > limit: + raise ValueError('past predicate boundary') + return value > 0 and value % 3 == 0 + + for stop in [0, 1, 7, 8, 9, 16, 17]: + limit = stop * 3 + for source_raises in [False, True]: + events = list[int]() + source = (raises_after_prefix(limit) if source_raises + else observed(events, limit + 2)) + result = sum(islice((value for value in source + if checked_predicate(value, limit)), stop)) + assert result == 3 * stop * (stop + 1) // 2 + if source_raises: + try: + next(source) + assert False + except ValueError as error: + assert error.message == 'past prefix' + else: + assert events == ([] if stop == 0 else list(range(limit + 1))) + events = list[int]() assert sum(islice(observed(events, 7), 3)) == 3 assert events == [0, 1, 2] diff --git a/test/stdlib/sort_test.codon b/test/stdlib/sort_test.codon index 7b6004edd..3c675cbaa 100644 --- a/test/stdlib/sort_test.codon +++ b/test/stdlib/sort_test.codon @@ -95,6 +95,82 @@ def test_standard_sort(): test_standard_sort() +def check_native_sort_type(T: type, lowest: T, highest: T): + for size in (0, 1, 2, 7, 15, 16, 17, 63, 64, 65, 127, 128, 129, 1024): + original = [T((index * 17) % 31 + 1) for index in range(size)] + if size > 0: + original[0] = highest + if size > 1: + original[-1] = lowest + expected = sorted(original, algorithm="tim") + actual = original.copy() + alias = actual + actual.sort() + assert actual is alias + assert alias == expected, f"{T.__name__}, size {size}: {alias} != {expected}" + actual.sort(reverse=True) + assert actual == expected[::-1] + assert sorted(value for value in original) == expected + keyed_expected = original.copy() + tim_sort_inplace(keyed_expected, lambda value: 0) + actual = original.copy() + actual.sort(key=lambda value: 0) + assert actual == keyed_expected + actual.sort(algorithm="pdq") + assert actual == expected + assert original == [highest if index == 0 else lowest if index == size - 1 + else T((index * 17) % 31 + 1) for index in range(size)] + + +def check_special_float_sort(T: type): + from math import copysign, isnan + for original in ([T(3), T(-0.0), T(0.0), T(1), T(-2)], + [T("nan"), T(3), T(-1), T("nan"), T(2)], + [T(3), T("nan"), T(-1), T("inf"), T("-inf")]): + for reverse in (False, True): + expected = sorted(original, reverse=reverse, algorithm="pdq") + actual = original.copy() + actual.sort(reverse=reverse) + assert len(actual) == len(expected) + for value, reference in zip(actual, expected): + if isnan(float(reference)): + assert isnan(float(value)) + else: + assert value == reference + assert copysign(1.0, float(value)) == copysign(1.0, float(reference)) + + +@test +def test_native_sort(): + from internal.sort import _try_native_sort + + check_native_sort_type(int, -9223372036854775808, 9223372036854775807) + check_native_sort_type(i16, i16(-32768), i16(32767)) + check_native_sort_type(u16, u16(0), u16(65535)) + check_native_sort_type(i32, i32(-2147483648), i32(2147483647)) + check_native_sort_type(u32, u32(0), u32(4294967295)) + check_native_sort_type(i64, i64(-9223372036854775808), i64(9223372036854775807)) + check_native_sort_type(u64, u64(0), u64(18446744073709551615)) + check_native_sort_type(float32, float32("-inf"), float32("inf")) + check_native_sort_type(float, float("-inf"), float("inf")) + check_native_sort_type(i8, i8(-128), i8(127)) + check_native_sort_type(u8, u8(0), u8(255)) + check_native_sort_type(i128, -(i128(1) << 100), i128(1) << 100) + check_native_sort_type(u128, u128(0), u128(1) << 100) + wide = [u128(1) << 127, (u128(1) << 64) + u128(1), + u128(1) << 64, ~u128(0), u128(0)] + native_sorted = _try_native_sort(wide._ptr, len(wide)) + assert native_sorted + assert wide == [u128(0), u128(1) << 64, (u128(1) << 64) + u128(1), + u128(1) << 127, ~u128(0)] + check_native_sort_type(bool, False, True) + check_special_float_sort(float32) + check_special_float_sort(float) + + +test_native_sort() + + @test def test_timsort_minrun(): from algorithms.timsort import _tim_sort From 2cc4bad54e4bb9186db2c602aa9a0d2e0fa82b9c Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 00:52:28 -0400 Subject: [PATCH 03/12] Update plugin API --- CMakeLists.txt | 14 +- codon/cir/llvm/optimize.cpp | 4 +- codon/compiler/jit.cpp | 4 + codon/dsl/dsl.h | 18 ++- codon/dsl/plugins.cpp | 54 +++++-- codon/dsl/plugins.h | 4 + docs/developers/extend.md | 43 ++++- test/dsl/library.cpp | 11 ++ test/dsl/plugins.cpp | 303 ++++++++++++++++++++++++++++++++++++ 9 files changed, 431 insertions(+), 24 deletions(-) create mode 100644 test/dsl/library.cpp create mode 100644 test/dsl/plugins.cpp diff --git a/CMakeLists.txt b/CMakeLists.txt index ce239496d..cedd9f765 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -546,8 +546,17 @@ set(CODON_TEST_CPPFILES test/cir/util/matching.cpp test/cir/value.cpp test/cir/var.cpp + test/dsl/plugins.cpp test/types.cpp) +add_library(codon_test_compiler_plugin SHARED EXCLUDE_FROM_ALL test/dsl/library.cpp) +target_compile_definitions(codon_test_compiler_plugin PRIVATE CODON_TEST_COMPILER_PLUGIN) +target_include_directories(codon_test_compiler_plugin PRIVATE + "${CMAKE_CURRENT_SOURCE_DIR}" ${LLVM_INCLUDE_DIRS} + "${peglib_SOURCE_DIR}") +target_link_libraries(codon_test_compiler_plugin PRIVATE codonc fmt) +add_library(codon_test_runtime_plugin SHARED EXCLUDE_FROM_ALL test/dsl/library.cpp) add_executable(codon_test ${CODON_TEST_CPPFILES}) +add_dependencies(codon_test codon codon_test_compiler_plugin codon_test_runtime_plugin) target_include_directories(codon_test PRIVATE test/cir "${gc_SOURCE_DIR}/include") target_link_libraries(codon_test PRIVATE fmt codonc codonrt gtest_main) if(ASAN) @@ -559,7 +568,10 @@ if(ASAN) "-fsanitize-recover=address") endif() target_compile_definitions(codon_test - PRIVATE TEST_DIR="${CMAKE_CURRENT_SOURCE_DIR}/test") + PRIVATE TEST_DIR="${CMAKE_CURRENT_SOURCE_DIR}/test" + TEST_CODON="$" + TEST_PLUGIN_COMPILER="$" + TEST_PLUGIN_RUNTIME="$") if(APPLE) set_target_properties( diff --git a/codon/cir/llvm/optimize.cpp b/codon/cir/llvm/optimize.cpp index dd0d29b00..d15a5030f 100644 --- a/codon/cir/llvm/optimize.cpp +++ b/codon/cir/llvm/optimize.cpp @@ -1490,14 +1490,14 @@ void runLLVMOptimizationPasses(llvm::Module *module, PluginManager *plugins, llvm::TargetLibraryInfoImpl tlii(moduleTriple); fam.registerPass([&] { return llvm::TargetLibraryAnalysis(tlii); }); + registerCodonLLVMOptimizationPasses(pb, plugins, options); + pb.registerModuleAnalyses(mam); pb.registerCGSCCAnalyses(cgam); pb.registerFunctionAnalyses(fam); pb.registerLoopAnalyses(lam); pb.crossRegisterProxies(lam, fam, cgam, mam); - registerCodonLLVMOptimizationPasses(pb, plugins, options); - if (options->debug) { llvm::ModulePassManager mpm = pb.buildO0DefaultPipeline(llvm::OptimizationLevel::O0); diff --git a/codon/compiler/jit.cpp b/codon/compiler/jit.cpp index 3affc2bd2..8694bbcf4 100644 --- a/codon/compiler/jit.cpp +++ b/codon/compiler/jit.cpp @@ -63,6 +63,8 @@ void collectExecutableStmts(ast::Stmt *s, ast::SuiteStmt *final) { } llvm::Error JIT::init(bool forgetful) { + if (auto error = compiler->getPluginManager()->loadRuntimeLibraries()) + return error; if (forgetful) { this->forgetful = true; auto fs = @@ -106,6 +108,8 @@ llvm::Error JIT::init(bool forgetful) { } llvm::Error JIT::compile(const ir::Func *input, llvm::orc::ResourceTrackerSP rt) { + if (auto error = compiler->getPluginManager()->loadRuntimeLibraries()) + return error; auto *module = compiler->getModule(); auto *pm = compiler->getPassManager(); auto *llvisitor = compiler->getLLVMVisitor(); diff --git a/codon/dsl/dsl.h b/codon/dsl/dsl.h index 68d5025fa..7e8eb7660 100644 --- a/codon/dsl/dsl.h +++ b/codon/dsl/dsl.h @@ -8,6 +8,7 @@ #include "codon/parser/cache.h" #include "llvm/Passes/PassBuilder.h" #include +#include #include #include @@ -32,10 +33,21 @@ class DSL { std::string supported; /// Plugin stdlib path std::string stdlibPath; - /// Plugin dynamic library path + /// Legacy library path, used for both compiler and runtime by default. std::string dylibPath; - /// Linker arguments (to replace "-l dylibPath" if present) + /// Linker arguments (to replace the default runtime library argument if present). std::vector linkArgs; + /// Compiler library: unset inherits dylibPath; empty disables loading. + std::optional compilerDylibPath; + /// Runtime library: unset inherits dylibPath; empty disables default linking. + std::optional runtimeDylibPath; + + const std::string &getCompilerDylibPath() const { + return compilerDylibPath ? *compilerDylibPath : dylibPath; + } + const std::string &getRuntimeDylibPath() const { + return runtimeDylibPath ? *runtimeDylibPath : dylibPath; + } }; using KeywordCallback = @@ -60,6 +72,8 @@ class DSL { virtual void addIRPasses(ir::transform::PassManager *pm, bool debug) {} /// Registers this DSL's LLVM passes with the given pass builder. + /// Called before analysis registration and pipeline construction, allowing + /// plugins to register analyses as well as passes. /// @param pb the pass builder to add the passes to /// @param debug true if compiling in debug mode virtual void addLLVMPasses(llvm::PassBuilder *pb, bool debug) {} diff --git a/codon/dsl/plugins.cpp b/codon/dsl/plugins.cpp index ad268fa32..85824a18e 100644 --- a/codon/dsl/plugins.cpp +++ b/codon/dsl/plugins.cpp @@ -49,15 +49,23 @@ llvm::Expected PluginManager::load(const std::string &path) { auto about = tml["about"]; auto library = tml["library"]; - std::string cppLib = library["cpp"].value_or(""); - std::string dylibPath; - if (!cppLib.empty()) { + auto libraryPath = [&](const std::string &name) -> std::string { + if (name.empty()) + return {}; llvm::SmallString<128> p = llvm::sys::path::parent_path(tomlPath); - llvm::sys::path::append(p, cppLib + "." + libExt); - dylibPath = p.str(); + llvm::sys::path::append(p, name + "." + libExt); + return std::string(p.str()); + }; + for (const char *key : {"compiler", "runtime"}) { + if (library[key] && !library[key].is_string()) + return pluginError(fmt::format("library.{} must be a string", key)); } + auto compilerLib = library["compiler"].value(); + auto runtimeLib = library["runtime"].value(); auto link = library["link"]; + if (!runtimeLib && link.is_boolean() && !link.value_or(true)) + runtimeLib = ""; std::vector linkArgs; if (auto arr = link.as_array()) { arr->for_each([&linkArgs](auto &&el) { @@ -82,14 +90,17 @@ llvm::Expected PluginManager::load(const std::string &path) { stdlibPath = p.str(); } - DSL::Info info = {about["name"].value_or(""), - about["description"].value_or(""), - about["version"].value_or(""), - about["url"].value_or(""), - about["supported"].value_or(""), - stdlibPath, - dylibPath, - linkArgs}; + DSL::Info info = { + about["name"].value_or(""), + about["description"].value_or(""), + about["version"].value_or(""), + about["url"].value_or(""), + about["supported"].value_or(""), + stdlibPath, + libraryPath(library["cpp"].value_or("")), + linkArgs, + compilerLib ? std::make_optional(libraryPath(*compilerLib)) : std::nullopt, + runtimeLib ? std::make_optional(libraryPath(*runtimeLib)) : std::nullopt}; bool versionOk = false; try { @@ -104,6 +115,7 @@ llvm::Expected PluginManager::load(const std::string &path) { return pluginError(fmt::format("unsupported version {} (supported: {})", CODON_VERSION, info.supported)); + const auto &dylibPath = info.getCompilerDylibPath(); if (!dylibPath.empty()) { std::string libLoadErrorMsg; auto handle = llvm::sys::DynamicLibrary::getPermanentLibrary(dylibPath.c_str(), @@ -127,4 +139,20 @@ llvm::Expected PluginManager::load(const std::string &path) { return plugins.back().get(); } +llvm::Error PluginManager::loadRuntimeLibraries() const { + for (const auto &plugin : plugins) { + const auto &runtimePath = plugin->info.getRuntimeDylibPath(); + if (runtimePath.empty() || runtimePath == plugin->info.getCompilerDylibPath() || + loadedRuntimeLibraries.count(runtimePath)) + continue; + std::string message; + if (llvm::sys::DynamicLibrary::LoadLibraryPermanently(runtimePath.c_str(), + &message)) + return llvm::make_error( + fmt::format("could not load runtime library '{}': {}", runtimePath, message)); + loadedRuntimeLibraries.insert(runtimePath); + } + return llvm::Error::success(); +} + } // namespace codon diff --git a/codon/dsl/plugins.h b/codon/dsl/plugins.h index 487a0e264..e64255b39 100644 --- a/codon/dsl/plugins.h +++ b/codon/dsl/plugins.h @@ -5,6 +5,7 @@ #include #include #include +#include #include #include "codon/cir/util/iterators.h" @@ -34,6 +35,7 @@ class PluginManager { std::string argv0; /// vector of loaded plugins std::vector> plugins; + mutable std::unordered_set loadedRuntimeLibraries; public: /// Constructs a plugin manager @@ -52,6 +54,8 @@ class PluginManager { /// @param path path to plugin directory containing "plugin.toml" file /// @return plugin pointer if successful, plugin error otherwise llvm::Expected load(const std::string &path); + + llvm::Error loadRuntimeLibraries() const; }; } // namespace codon diff --git a/docs/developers/extend.md b/docs/developers/extend.md index 085dc3b53..54815de18 100644 --- a/docs/developers/extend.md +++ b/docs/developers/extend.md @@ -12,17 +12,46 @@ called `plugin.toml`. The following fields are supported: - `about.version`: Plugin version, using [semantic versioning](https://semver.org) - `about.url`: Plugin URL - `about.supported`: Supported Codon versions, using semantic versioning ranges -- `library.cpp`: Shared library to be loaded upon loading the plugin, which includes - the plugin implementation (see below) and any necessary runtime functions. The - library extension (i.e. `.so` or `.dylib`) will be added automatically, and should - not be included. +- `library.cpp`: Legacy shared library, used for both the compiler implementation + and runtime functions unless overridden below. Existing configurations retain + their behavior. +- `library.compiler`: Shared library containing the compiler plugin implementation + and its `load()` entry point. It is loaded into the compiler and is linked into + generated programs only if also selected as the runtime library. An omitted value + inherits `library.cpp`; an empty string disables compiler-library loading. +- `library.runtime`: Shared library containing runtime functions used by generated + code. It is linked into executables and shared libraries, and loaded for execution + by `codon run` or the JIT. It does not need a `load()` entry point and is not loaded + during ordinary compilation. An omitted value inherits `library.cpp`; an empty + string disables automatic runtime linking and runtime-library loading. + All three library paths are relative to the manifest directory. Their extensions + (`.so` or `.dylib`) are added automatically and should not be included. - `library.codon`: Standard library code that should be included with the plugin. It is recommended to put this code in directory `stdlib/`, whereupon the value of this parameter would be `"stdlib"`. - `library.link`: Libraries to be linked when compiling to an executable. The string `{root}` will be replaced with the path to the TOML configuration file. For example, a value similar to `{root}/build/libmyplugin.a` might be used, assuming the plugin - also builds a static library containing necessary runtime functions. + also builds a static library containing necessary runtime functions. A string or + array of strings overrides the default runtime library argument. Explicit link + arguments remain usable even without a runtime library. For compatibility, + `library.link = false` means `library.runtime = ""` when `library.runtime` is omitted. + Search paths come from the runtime library, not the compiler library. + +For separate compiler and runtime libraries: + +```toml +[library] +compiler = "build/libmyplugin_compiler" +runtime = "build/libmyplugin_runtime" +``` + +For a compiler-only plugin, set `runtime = ""`. A runtime-only plugin can omit both +`compiler` and `cpp`. In C++, `DSL::Info::compilerDylibPath` and `runtimeDylibPath` +are optional: unset values inherit `dylibPath`, while empty strings disable the +corresponding library. This preserves legacy configuration and normal source +initializers, not the binary layout of `DSL::Info`; plugins using that metadata +must be rebuilt against the matching headers. Here is an example configuration file for the validate pass shown in the [Codon IR docs](ir.md#bidirectionality): @@ -79,4 +108,6 @@ can be found [on GitHub](https://github.com/exaloop/example-codon-plugin). Plugins can add new LLVM passes by overriding the `void addLLVMPasses(llvm::PassBuilder *pb, bool debug)` method of the `codon::DSL` class. Refer to the [`llvm::PassBuilder` docs](https://llvm.org/doxygen/classllvm_1_1PassBuilder.html) for details on adding -passes. +passes. This hook runs before analysis registration and pipeline construction, so +plugins can also use `registerAnalysisRegistrationCallback` to register analyses. The hook is called for +each LLVM pipeline construction. Callbacks should not modify global LLVM command-line options. diff --git a/test/dsl/library.cpp b/test/dsl/library.cpp new file mode 100644 index 000000000..e0ec91bd7 --- /dev/null +++ b/test/dsl/library.cpp @@ -0,0 +1,11 @@ +#include + +#ifdef CODON_TEST_COMPILER_PLUGIN +#include "codon/dsl/dsl.h" + +extern "C" std::unique_ptr load() { return std::make_unique(); } + +extern "C" int64_t codon_test_compiler_value() { return 17; } +#else +extern "C" int64_t codon_test_runtime_value() { return 42; } +#endif diff --git a/test/dsl/plugins.cpp b/test/dsl/plugins.cpp new file mode 100644 index 000000000..6cdc8e973 --- /dev/null +++ b/test/dsl/plugins.cpp @@ -0,0 +1,303 @@ +#include "codon/dsl/plugins.h" +#include "codon/compiler/jit.h" +#include "llvm/ADT/SmallString.h" +#include "llvm/Support/FileSystem.h" +#include "llvm/Support/MemoryBuffer.h" +#include "llvm/Support/Path.h" +#include "llvm/Support/Program.h" +#include "llvm/Support/raw_ostream.h" +#include "gtest/gtest.h" + +namespace codon { +namespace { + +TEST(PluginInfoTest, LegacyLibraryDefaultsToCompilerAndRuntime) { + DSL::Info info = {"name", "description", "1.0.0", "url", + "*", "stdlib", "legacy", {"-lextra"}}; + EXPECT_EQ(info.getCompilerDylibPath(), "legacy"); + EXPECT_EQ(info.getRuntimeDylibPath(), "legacy"); + + info.compilerDylibPath = "compiler"; + EXPECT_EQ(info.getCompilerDylibPath(), "compiler"); + EXPECT_EQ(info.getRuntimeDylibPath(), "legacy"); + info.runtimeDylibPath = "runtime"; + EXPECT_EQ(info.getRuntimeDylibPath(), "runtime"); + + info.runtimeDylibPath = ""; + EXPECT_TRUE(info.getRuntimeDylibPath().empty()); + EXPECT_EQ(info.getCompilerDylibPath(), "compiler"); + info.compilerDylibPath = ""; + EXPECT_TRUE(info.getCompilerDylibPath().empty()); + info.compilerDylibPath.reset(); + info.runtimeDylibPath.reset(); + EXPECT_EQ(info.getCompilerDylibPath(), "legacy"); + EXPECT_EQ(info.getRuntimeDylibPath(), "legacy"); +} + +class PluginLibrariesTest : public ::testing::Test { +protected: + llvm::SmallString<128> directory; + PluginManager manager{""}; + std::string compilerStem; + std::string runtimeStem; + + void SetUp() override { + auto error = + llvm::sys::fs::createUniqueDirectory("codon-plugin-libraries", directory); + ASSERT_FALSE(error) << error.message(); + compilerStem = "compiler/" + llvm::sys::path::stem(TEST_PLUGIN_COMPILER).str(); + runtimeStem = "runtime/" + llvm::sys::path::stem(TEST_PLUGIN_RUNTIME).str(); + for (const auto &entry : {std::make_pair(TEST_PLUGIN_COMPILER, compilerStem), + std::make_pair(TEST_PLUGIN_RUNTIME, runtimeStem)}) { + auto destination = libraryPath(entry.second); + error = + llvm::sys::fs::create_directories(llvm::sys::path::parent_path(destination)); + ASSERT_FALSE(error) << error.message(); + error = llvm::sys::fs::copy_file(entry.first, destination); + ASSERT_FALSE(error) << error.message(); + } + } + + void TearDown() override { llvm::sys::fs::remove_directories(directory); } + + void manifest(const std::string &libraries) { + llvm::SmallString<128> filename(directory); + llvm::sys::path::append(filename, "plugin.toml"); + std::error_code error; + llvm::raw_fd_ostream output(filename, error); + ASSERT_FALSE(error) << error.message(); + output << "[about]\nname = \"test\"\nsupported = \">=0.0.0\"\n[library]\n" + << libraries << '\n'; + } + + std::string libraryPath(const std::string &name) { + llvm::SmallString<128> filename(directory); +#ifdef __APPLE__ + llvm::sys::path::append(filename, name + ".dylib"); +#else + llvm::sys::path::append(filename, name + ".so"); +#endif + return std::string(filename.str()); + } + + std::string execute(const std::vector &arguments) { + std::vector references(arguments.begin(), arguments.end()); + std::string outputPath = std::string(directory.str()) + "/stdout.txt"; + std::string errorPath = std::string(directory.str()) + "/stderr.txt"; + std::vector> redirects = { + std::nullopt, llvm::StringRef(outputPath), llvm::StringRef(errorPath)}; + std::string message; + int result = llvm::sys::ExecuteAndWait(arguments.front(), references, std::nullopt, + redirects, 60, 0, &message); + auto output = llvm::MemoryBuffer::getFile(outputPath); + auto errors = llvm::MemoryBuffer::getFile(errorPath); + EXPECT_EQ(result, 0) << message << (errors ? (*errors)->getBuffer().str() : ""); + return output ? (*output)->getBuffer().str() : ""; + } + + void checkExecution(const std::string &libraries, const std::string &code, + const std::string &expected, bool linksCompiler, + bool linksRuntime) { + manifest(libraries); + std::string root(directory.str()); + std::string source = root + "/program.codon"; + std::string executable = root + "/program"; + std::error_code error; + { + llvm::raw_fd_ostream output(source, error); + ASSERT_FALSE(error) << error.message(); + output << code; + } + execute( + {TEST_CODON, "build", "-release", "-plugin", root, "-o", executable, source}); + ASSERT_FALSE(HasFailure()); + EXPECT_EQ(execute({executable}), expected); + EXPECT_EQ(execute({TEST_CODON, "run", "-release", "-plugin", root, source}), + expected); +#ifdef __APPLE__ + auto dependencies = execute({"/usr/bin/otool", "-L", executable}); + EXPECT_EQ( + dependencies.find(llvm::sys::path::filename(TEST_PLUGIN_COMPILER).str()) != + std::string::npos, + linksCompiler); + EXPECT_EQ(dependencies.find(llvm::sys::path::filename(TEST_PLUGIN_RUNTIME).str()) != + std::string::npos, + linksRuntime); + if (!linksCompiler) { + auto commands = execute({"/usr/bin/otool", "-l", executable}); + EXPECT_EQ(commands.find(root + "/compiler"), std::string::npos); + } +#endif + } +}; + +TEST_F(PluginLibrariesTest, RuntimeOnlyLibraryIsNotLoadedIntoCompiler) { + manifest("runtime = \"missing-runtime\""); + auto result = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(result)) << llvm::toString(result.takeError()); + EXPECT_TRUE((*result)->info.getCompilerDylibPath().empty()); + EXPECT_EQ((*result)->info.getRuntimeDylibPath(), libraryPath("missing-runtime")); + auto error = manager.loadRuntimeLibraries(); + ASSERT_TRUE(bool(error)); + EXPECT_NE(llvm::toString(std::move(error)).find("missing-runtime"), + std::string::npos); +} + +TEST_F(PluginLibrariesTest, EmptyCompilerOverridesLegacyAndRuntimeInheritsIt) { + manifest("cpp = \"missing-legacy\"\ncompiler = \"\""); + auto result = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(result)) << llvm::toString(result.takeError()); + EXPECT_EQ((*result)->info.dylibPath, libraryPath("missing-legacy")); + EXPECT_TRUE((*result)->info.getCompilerDylibPath().empty()); + EXPECT_EQ((*result)->info.getRuntimeDylibPath(), (*result)->info.dylibPath); +} + +TEST_F(PluginLibrariesTest, RuntimeLibraryLoadFailuresAreRetried) { + manifest("runtime = \"retry-runtime\""); + auto result = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(result)) << llvm::toString(result.takeError()); + for (int attempt = 0; attempt < 2; ++attempt) { + auto error = manager.loadRuntimeLibraries(); + ASSERT_TRUE(bool(error)); + EXPECT_NE(llvm::toString(std::move(error)).find("retry-runtime"), + std::string::npos); + } + auto copyError = + llvm::sys::fs::copy_file(TEST_PLUGIN_RUNTIME, libraryPath("retry-runtime")); + ASSERT_FALSE(copyError) << copyError.message(); + for (int attempt = 0; attempt < 2; ++attempt) { + auto error = manager.loadRuntimeLibraries(); + ASSERT_FALSE(bool(error)) << llvm::toString(std::move(error)); + } +} + +TEST_F(PluginLibrariesTest, LoadsNewRuntimeLibrariesAfterPreviousSuccess) { + manifest("runtime = \"" + runtimeStem + "\""); + auto first = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(first)) << llvm::toString(first.takeError()); + auto error = manager.loadRuntimeLibraries(); + ASSERT_FALSE(bool(error)) << llvm::toString(std::move(error)); + + manifest("runtime = \"late-runtime\""); + auto second = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(second)) << llvm::toString(second.takeError()); + error = manager.loadRuntimeLibraries(); + ASSERT_TRUE(bool(error)); + EXPECT_NE(llvm::toString(std::move(error)).find("late-runtime"), std::string::npos); + + auto copyError = + llvm::sys::fs::copy_file(TEST_PLUGIN_RUNTIME, libraryPath("late-runtime")); + ASSERT_FALSE(copyError) << copyError.message(); + error = manager.loadRuntimeLibraries(); + ASSERT_FALSE(bool(error)) << llvm::toString(std::move(error)); +} + +TEST_F(PluginLibrariesTest, EmptyRuntimeOverridesLegacy) { + manifest("cpp = \"missing-legacy\"\ncompiler = \"\"\nruntime = \"\""); + auto result = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(result)) << llvm::toString(result.takeError()); + EXPECT_TRUE((*result)->info.getRuntimeDylibPath().empty()); + auto error = manager.loadRuntimeLibraries(); + EXPECT_FALSE(bool(error)) << llvm::toString(std::move(error)); +} + +TEST_F(PluginLibrariesTest, CompilerOverrideSelectsTheLoadedLibrary) { + manifest("cpp = \"missing-legacy\"\ncompiler = \"missing-compiler\"\nruntime = \"\""); + auto result = manager.load(std::string(directory.str())); + ASSERT_FALSE(bool(result)); + auto message = llvm::toString(result.takeError()); + EXPECT_NE(message.find("missing-compiler"), std::string::npos); + EXPECT_EQ(message.find("missing-legacy"), std::string::npos); +} + +TEST_F(PluginLibrariesTest, EmptyRuntimeDoesNotDisableLegacyCompilerLoading) { + manifest("cpp = \"missing-legacy\"\nruntime = \"\""); + auto result = manager.load(std::string(directory.str())); + ASSERT_FALSE(bool(result)); + EXPECT_NE(llvm::toString(result.takeError()).find("missing-legacy"), + std::string::npos); +} + +TEST_F(PluginLibrariesTest, LegacyLinkFalseMeansEmptyRuntime) { + manifest("cpp = \"missing-legacy\"\ncompiler = \"\"\nlink = false"); + auto result = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(result)) << llvm::toString(result.takeError()); + EXPECT_TRUE((*result)->info.getRuntimeDylibPath().empty()); +} + +TEST_F(PluginLibrariesTest, ExplicitLinkArgumentsRemainIndependent) { + manifest("runtime = \"\"\nlink = [\"-lstandalone\", \"{root}/custom.a\"]"); + auto result = manager.load(std::string(directory.str())); + ASSERT_TRUE(bool(result)) << llvm::toString(result.takeError()); + ASSERT_EQ((*result)->info.linkArgs.size(), 2); + EXPECT_EQ((*result)->info.linkArgs[0], "-lstandalone"); + EXPECT_EQ((*result)->info.linkArgs[1], std::string(directory.str()) + "/custom.a"); +} + +TEST_F(PluginLibrariesTest, RejectsNonStringLibraryPaths) { + manifest("runtime = false"); + auto result = manager.load(std::string(directory.str())); + ASSERT_FALSE(bool(result)); + EXPECT_NE(llvm::toString(result.takeError()).find("library.runtime must be a string"), + std::string::npos); +} + +TEST_F(PluginLibrariesTest, ExecutesWithSeparateCompilerAndRuntimeLibraries) { + checkExecution("compiler = \"" + compilerStem + "\"\nruntime = \"" + runtimeStem + + "\"", + "from C import codon_test_runtime_value() -> int\n" + "print(codon_test_runtime_value())\n", + "42\n", false, true); +} + +TEST_F(PluginLibrariesTest, ExecutesWithLegacyCombinedLibrary) { + checkExecution("cpp = \"" + compilerStem + "\"", + "from C import codon_test_compiler_value() -> int\n" + "print(codon_test_compiler_value())\n", + "17\n", true, false); +} + +TEST_F(PluginLibrariesTest, ExecutesWithOnlyRuntimeLibrary) { + checkExecution("runtime = \"" + runtimeStem + "\"", + "from C import codon_test_runtime_value() -> int\n" + "print(codon_test_runtime_value())\n", + "42\n", false, true); +} + +TEST_F(PluginLibrariesTest, ExecutesWithCompilerOnlyAndEmptyRuntime) { + checkExecution("cpp = \"" + compilerStem + "\"\nruntime = \"\"", "print(7)\n", "7\n", + false, false); +} + +TEST_F(PluginLibrariesTest, ExecutesWithExplicitRuntimeLinkArguments) { + checkExecution("compiler = \"" + compilerStem + "\"\nruntime = \"" + runtimeStem + + "\"\nlink = [\"{root}/" + runtimeStem + +#ifdef __APPLE__ + ".dylib\"]", +#else + ".so\"]", +#endif + "from C import codon_test_runtime_value() -> int\n" + "print(codon_test_runtime_value())\n", + "42\n", false, true); +} + +TEST_F(PluginLibrariesTest, InteractiveJITLoadsRuntimeForLatePlugin) { + Options options; + options.argv0 = TEST_CODON; + jit::JIT instance(options, "", std::string(TEST_DIR) + "/../stdlib"); + auto error = instance.init(); + ASSERT_FALSE(bool(error)) << llvm::toString(std::move(error)); + manifest("runtime = \"" + runtimeStem + "\""); + auto plugin = + instance.getCompiler()->getPluginManager()->load(std::string(directory.str())); + ASSERT_TRUE(bool(plugin)) << llvm::toString(plugin.takeError()); + for (int cell = 0; cell < 2; ++cell) { + auto result = instance.execute("from C import codon_test_runtime_value() -> int\n" + "assert codon_test_runtime_value() == 42\n"); + ASSERT_TRUE(bool(result)) << llvm::toString(result.takeError()); + } +} + +} // namespace +} // namespace codon From c2dfb149c5d380d84cedca92820ba02bbd70e152 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 00:59:00 -0400 Subject: [PATCH 04/12] Remove comment --- codon/cir/transform/manager.cpp | 2 -- 1 file changed, 2 deletions(-) diff --git a/codon/cir/transform/manager.cpp b/codon/cir/transform/manager.cpp index 49e0a062f..146fdf334 100644 --- a/codon/cir/transform/manager.cpp +++ b/codon/cir/transform/manager.cpp @@ -212,8 +212,6 @@ void PassManager::registerStandardPasses() { registerPass(std::make_unique(numpyKey, seKey2), /*insertBefore=*/"", {numpyKey, seKey2}, {seKey1, rdKey, cfgKey, globalKey, capKey}); - // Expose whole producer/consumer loops before lowering. LLVM's suspension-aware - // unroll guard alone does not eliminate nested scan/consumer loop structure. registerPass(std::make_unique(), /*insertBefore=*/"", {}, {seKey1, seKey2, rdKey, cfgKey, globalKey, capKey}); From 11a59a5f737a1afdafe8857c7a4a255cf1b7b539 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 01:07:11 -0400 Subject: [PATCH 05/12] Add comments --- codon/cir/transform/pythonic/generator.cpp | 40 ++++++++++++++++++++++ 1 file changed, 40 insertions(+) diff --git a/codon/cir/transform/pythonic/generator.cpp b/codon/cir/transform/pythonic/generator.cpp index 555c00b48..225d4c243 100644 --- a/codon/cir/transform/pythonic/generator.cpp +++ b/codon/cir/transform/pythonic/generator.cpp @@ -14,6 +14,9 @@ namespace ir { namespace transform { namespace pythonic { namespace { +// Two complementary rewrites avoid the generator protocol: turn direct +// sum/any/all calls into ordinary functions, and splice a producer's body into +// its sole for-loop consumer. The helpers below first build the reducer functions. bool isSum(Func *f) { return f && f->getName().rfind(ast::getMangledFunc("std.internal.builtin", "sum"), 0) == 0; @@ -100,6 +103,7 @@ struct GeneratorAnyAllTransformer : public util::Operator { } void handle(ReturnInstr *v) override { + // Leave the short-circuit returns inserted at yield sites unchanged. if (saw(v)) return; auto *M = v->getModule(); @@ -115,6 +119,8 @@ struct GeneratorAnyAllTransformer : public util::Operator { void handle(YieldInInstr *v) override { valid = false; } }; +// Clone rather than mutate the producer, which may still be used as a generator +// elsewhere. The wrapper takes the same arguments plus the initial accumulator. Func *genToSum(BodiedFunc *gen, Type *startType, Type *outType) { if (!gen || !gen->isGenerator()) return nullptr; @@ -169,6 +175,8 @@ Func *genToSum(BodiedFunc *gen, Type *startType, Type *outType) { return fn; } +// Use the same cloning scheme for short-circuit reducers. Original returns and +// fallthrough mean exhaustion, whose answer is false for any and true for all. Func *genToAnyAll(BodiedFunc *gen, bool any) { if (!gen || !gen->isGenerator()) return nullptr; @@ -210,6 +218,9 @@ Func *genToAnyAll(BodiedFunc *gen, bool any) { } // namespace namespace { +// Restrict fusion to small synchronous bodies. Extra suspension protocols, +// exception regions, and exposed local storage need lifetime/control-flow +// handling that the substitution below does not implement. struct FusionVerifier : public util::Operator { int nodes = 0; int yields = 0; @@ -230,6 +241,8 @@ struct FusionVerifier : public util::Operator { } }; +// Moving a consumer body into the producer changes its enclosing loops. Reject +// user break/continue, but allow breaks targeting compiler-created exit wrappers. struct ConsumerVerifier : public util::Operator { const std::unordered_set &wrappers; bool valid = true; @@ -241,6 +254,9 @@ struct ConsumerVerifier : public util::Operator { void handle(ContinueInstr *) override { valid = false; } }; +// A named iterator is eligible only when one assignment feeds one later read, +// with no exposed handle and matching enclosing control-flow contexts. This +// prevents fusion from changing the lifetime of shared or repeatedly used iterators. struct IteratorUseVerifier : public util::Operator { Var *iterator; ForFlow *consumer; @@ -258,6 +274,8 @@ struct IteratorUseVerifier : public util::Operator { const std::unordered_set &wrappers) : iterator(iterator), consumer(consumer), wrappers(wrappers) {} std::vector enclosingLoops() { + // Ignore exit wrappers introduced by this pass. Include branch identities + // so creation in one if/try arm cannot be paired with consumption in another. std::vector result; for (auto position = parent_begin(); position != parent_end(); ++position) { auto *node = cast(*position); @@ -306,6 +324,9 @@ struct IteratorUseVerifier : public util::Operator { } }; +// Traverse the producer bottom-up: the consumer body inserted at a yield must +// not itself be rewritten. In particular, its returns still exit the caller, +// while the producer's returns become exits from the synthetic wrapper loop. struct FusionTransformer : public util::Operator { ForFlow *consumer; WhileFlow *exit; @@ -321,6 +342,7 @@ struct FusionTransformer : public util::Operator { void handle(ReturnInstr *value) override { auto *module = value->getModule(); auto *replacement = module->Nr(); + // Exhaustion discards the producer's return value, but not its side effects. if (value->getValue()) replacement->push_back(value->getValue()); replacement->push_back(module->Nr(exit)); @@ -331,6 +353,9 @@ struct FusionTransformer : public util::Operator { const std::string GeneratorLoopFusion::KEY = "core-pythonic-generator-loop-fusion"; +// Reducer specialization can hide another producer/consumer pair behind this +// compiler-generated call. Inline only that wrapper to expose further fusion, +// charging the same per-caller growth budget used for producer substitution. void GeneratorLoopFusion::handle(CallInstr *call) { auto *parent = cast(getParentFunc()); auto *function = cast(util::getFunc(call->getCallee())); @@ -345,6 +370,8 @@ void GeneratorLoopFusion::handle(CallInstr *call) { return; for (auto *variable : inlined.newVars) parent->push_back(variable); + // The inliner implements early returns with a one-shot loop. Remember it so + // later fusion accepts its targeted breaks and ignores it as a lifetime boundary. if (auto *expression = cast(inlined.result)) if (auto *body = cast(expression->getFlow())) if (auto *wrapper = cast(body->back())) @@ -359,6 +386,9 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { auto *parent = cast(getParentFunc()); if (!parent) return; + // Unpack a single-use tuple that forwards an iterator. Capture every element + // at the original assignment, in order, even if only one feeds this loop; + // dropping the other elements could discard effects. Then retry the direct case. if (auto *extract = cast(loop->getIter())) { auto *variable = util::getVar(extract->getVal()); if (!variable || variable->isGlobal()) @@ -388,6 +418,8 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { handle(loop); return; } + // Accept either an immediate generator call or a uniquely owned local handle. + // Retain the latter's assignment so argument evaluation stays at creation time. auto *call = cast(loop->getIter()); AssignInstr *creation = nullptr; if (auto *iterator = util::getVar(loop->getIter())) { @@ -411,6 +443,8 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { verifier.process(generator->getBody()); ConsumerVerifier consumerVerifier(wrappers); consumerVerifier.process(loop->getBody()); + // One yield site lets us move, rather than duplicate, the consumer body. + // The cumulative budget bounds growth when nested generators are fused in turn. if (!verifier.valid || verifier.yields != 1 || !consumerVerifier.valid || growth[parent->getId()] + verifier.nodes > 1024) return; @@ -420,6 +454,8 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { auto *setup = module->Nr(); std::unordered_map arguments; auto argument = generator->arg_begin(); + // Snapshot arguments once, left-to-right. For a named handle this setup replaces + // its creation, not its use, preserving mutations and exceptions between the two. for (auto *value : *call) arguments.emplace((*argument++)->getId(), util::makeVar(value, setup, parent)); if (creation) { @@ -428,6 +464,8 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { } util::CloneVisitor clone(module); auto *body = cast(clone.clone(generator->getBody(), parent, arguments)); + // A one-shot loop gives producer returns a common exit without returning from + // the caller. Reaching the end also exits; only the producer's own loops repeat. auto *wrapper = module->Nr(); auto *exit = module->Nr(module->getBool(true), wrapper); wrappers.insert(exit->getId()); @@ -442,6 +480,8 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { const std::string GeneratorArgumentOptimization::KEY = "core-pythonic-generator-argument-opt"; +// Specialize only reducers applied directly to a generator call, where its +// construction arguments can be forwarded to the wrapper without using a handle. void GeneratorArgumentOptimization::handle(CallInstr *v) { auto *M = v->getModule(); auto *func = util::getFunc(v->getCallee()); From 07a7e508e048a0da8ffb63215d026e6d2597bbf8 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 01:26:04 -0400 Subject: [PATCH 06/12] Update native sort dispatch --- stdlib/internal/sort.codon | 10 +--------- test/stdlib/sort_test.codon | 18 ++++++++++-------- 2 files changed, 11 insertions(+), 17 deletions(-) diff --git a/stdlib/internal/sort.codon b/stdlib/internal/sort.codon index c996274a2..8c49f35e4 100644 --- a/stdlib/internal/sort.codon +++ b/stdlib/internal/sort.codon @@ -39,14 +39,6 @@ def _try_native_sort(start: Ptr[T], n: int, T: type) -> bool: return False return True -def _try_native_sort_list(start: Ptr[T], n: int, T: type) -> bool: - if T is float or T is float32: - for index in range(n): - value = start[index] - if value != value or value == 0: - return False - return _try_native_sort(start, n) - def sorted( v: Generator[T], key=Optional[int](), @@ -107,7 +99,7 @@ class List: if isinstance(key, Optional): if algorithm == "auto": if self: - if _try_native_sort_list(self._ptr, self._len): + if _try_native_sort(self._ptr, self._len): pass elif _is_pdq_compatible(self[0]): pdq_sort_inplace(self, lambda x: x) diff --git a/test/stdlib/sort_test.codon b/test/stdlib/sort_test.codon index 3c675cbaa..7498e040b 100644 --- a/test/stdlib/sort_test.codon +++ b/test/stdlib/sort_test.codon @@ -125,19 +125,21 @@ def check_native_sort_type(T: type, lowest: T, highest: T): def check_special_float_sort(T: type): from math import copysign, isnan for original in ([T(3), T(-0.0), T(0.0), T(1), T(-2)], + [T(-0.0), T(0.0), T(0.0), T(-0.0)], + [T("nan"), T("nan")], [T("nan"), T(3), T(-1), T("nan"), T(2)], [T(3), T("nan"), T(-1), T("inf"), T("-inf")]): for reverse in (False, True): - expected = sorted(original, reverse=reverse, algorithm="pdq") + expected = sorted((value for value in original if not isnan(float(value))), + reverse=reverse, algorithm="pdq") actual = original.copy() actual.sort(reverse=reverse) - assert len(actual) == len(expected) - for value, reference in zip(actual, expected): - if isnan(float(reference)): - assert isnan(float(value)) - else: - assert value == reference - assert copysign(1.0, float(value)) == copysign(1.0, float(reference)) + assert len(actual) == len(original) + assert [value for value in actual if not isnan(float(value))] == expected + assert sum(1 for value in actual + if value == 0 and copysign(1.0, float(value)) < 0) == sum( + 1 for value in original + if value == 0 and copysign(1.0, float(value)) < 0) @test From 9d36aac5f75acf9b308e43631cb5ed0bfea3f204 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 01:35:22 -0400 Subject: [PATCH 07/12] Add more comments --- codon/cir/transform/pythonic/generator.cpp | 141 +++++++++++++++++++++ 1 file changed, 141 insertions(+) diff --git a/codon/cir/transform/pythonic/generator.cpp b/codon/cir/transform/pythonic/generator.cpp index 225d4c243..55bd490a5 100644 --- a/codon/cir/transform/pythonic/generator.cpp +++ b/codon/cir/transform/pythonic/generator.cpp @@ -14,9 +14,12 @@ namespace ir { namespace transform { namespace pythonic { namespace { + // Two complementary rewrites avoid the generator protocol: turn direct // sum/any/all calls into ordinary functions, and splice a producer's body into // its sole for-loop consumer. The helpers below first build the reducer functions. +// The sketches below use fresh local names and omit type conversions and +// eligibility checks; they illustrate successful rewrites, not arbitrary generators. bool isSum(Func *f) { return f && f->getName().rfind(ast::getMangledFunc("std.internal.builtin", "sum"), 0) == 0; @@ -106,6 +109,7 @@ struct GeneratorAnyAllTransformer : public util::Operator { // Leave the short-circuit returns inserted at yield sites unchanged. if (saw(v)) return; + auto *M = v->getModule(); auto *newReturn = M->Nr(M->getBool(!any)); see(newReturn); @@ -121,6 +125,27 @@ struct GeneratorAnyAllTransformer : public util::Operator { // Clone rather than mutate the producer, which may still be used as a generator // elsewhere. The wrapper takes the same arguments plus the initial accumulator. +// For example: +// +// def gen(items): +// for item in items: +// if keep(item): +// yield compute(item) +// result = sum(gen(items), start) +// +// becomes: +// +// def sum_gen(items, start): +// total = start +// for item in items: +// if keep(item): +// total = total + compute(item) +// return total +// result = sum_gen(items, start) +// +// There is no generator object at this call site: each yielded value contributes +// directly to the result. An early producer return instead returns total, after +// evaluating any return expression for its effects. Func *genToSum(BodiedFunc *gen, Type *startType, Type *outType) { if (!gen || !gen->isGenerator()) return nullptr; @@ -177,6 +202,24 @@ Func *genToSum(BodiedFunc *gen, Type *startType, Type *outType) { // Use the same cloning scheme for short-circuit reducers. Original returns and // fallthrough mean exhaustion, whose answer is false for any and true for all. +// For the producer above, any(gen(items)) becomes any_gen(items): +// +// def any_gen(items): +// for item in items: +// if keep(item) and bool(compute(item)): +// return True +// return False +// +// Likewise, all(gen(items)) becomes all_gen(items): +// +// def all_gen(items): +// for item in items: +// if keep(item) and not bool(compute(item)): +// return False +// return True +// +// Only the needed prefix of the producer runs; finding the answer skips the +// remaining producer code rather than first collecting all yielded values. Func *genToAnyAll(BodiedFunc *gen, bool any) { if (!gen || !gen->isGenerator()) return nullptr; @@ -218,6 +261,7 @@ Func *genToAnyAll(BodiedFunc *gen, bool any) { } // namespace namespace { + // Restrict fusion to small synchronous bodies. Extra suspension protocols, // exception regions, and exposed local storage need lifetime/control-flow // handling that the substitution below does not implement. @@ -227,15 +271,18 @@ struct FusionVerifier : public util::Operator { bool valid = true; void preHook(Node *) override { valid &= ++nodes <= 256; } + void handle(YieldInstr *value) override { valid &= value->getValue() && !value->isFinal(); ++yields; } + void handle(YieldInInstr *) override { valid = false; } void handle(AwaitInstr *) override { valid = false; } void handle(TryCatchFlow *) override { valid = false; } void handle(PointerValue *value) override { valid &= value->getVar()->isGlobal(); } void handle(StackAllocInstr *) override { valid = false; } + void handle(ForFlow *loop) override { valid &= !loop->isParallel() && !loop->isAsync(); } @@ -246,11 +293,14 @@ struct FusionVerifier : public util::Operator { struct ConsumerVerifier : public util::Operator { const std::unordered_set &wrappers; bool valid = true; + explicit ConsumerVerifier(const std::unordered_set &wrappers) : wrappers(wrappers) {} + void handle(BreakInstr *value) override { valid &= value->getLoop() && wrappers.count(value->getLoop()->getId()); } + void handle(ContinueInstr *) override { valid = false; } }; @@ -270,9 +320,11 @@ struct IteratorUseVerifier : public util::Operator { bool addressTaken = false; std::vector creationLoops; std::vector consumptionLoops; + IteratorUseVerifier(Var *iterator, ForFlow *consumer, const std::unordered_set &wrappers) : iterator(iterator), consumer(consumer), wrappers(wrappers) {} + std::vector enclosingLoops() { // Ignore exit wrappers introduced by this pass. Include branch identities // so creation in one if/try arm cannot be paired with consumption in another. @@ -291,15 +343,18 @@ struct IteratorUseVerifier : public util::Operator { } return result; } + void preHook(Node *node) override { ++position; auto *value = cast(node); if (!value || isA(value) || isA(value) || isA(value)) return; + for (auto *variable : value->getUsedVariables()) addressTaken |= variable->getId() == iterator->getId(); } + void handle(VarValue *value) override { if (value->getVar()->getId() == iterator->getId()) { ++reads; @@ -307,9 +362,11 @@ struct IteratorUseVerifier : public util::Operator { consumptionLoops = enclosingLoops(); } } + void handle(PointerValue *value) override { addressTaken |= value->getVar()->getId() == iterator->getId(); } + void handle(AssignInstr *value) override { if (value->getLhs()->getId() == iterator->getId()) { assignment = value; @@ -318,6 +375,7 @@ struct IteratorUseVerifier : public util::Operator { creationLoops = enclosingLoops(); } } + bool valid() const { return reads == 1 && writes == 1 && !addressTaken && creationPosition < consumptionPosition && creationLoops == consumptionLoops; @@ -330,6 +388,7 @@ struct IteratorUseVerifier : public util::Operator { struct FusionTransformer : public util::Operator { ForFlow *consumer; WhileFlow *exit; + FusionTransformer(ForFlow *consumer, WhileFlow *exit) : util::Operator(true), consumer(consumer), exit(exit) {} @@ -339,9 +398,11 @@ struct FusionTransformer : public util::Operator { util::series(module->Nr(consumer->getVar(), value->getValue()), consumer->getBody())); } + void handle(ReturnInstr *value) override { auto *module = value->getModule(); auto *replacement = module->Nr(); + // Exhaustion discards the producer's return value, but not its side effects. if (value->getValue()) replacement->push_back(value->getValue()); @@ -356,55 +417,123 @@ const std::string GeneratorLoopFusion::KEY = "core-pythonic-generator-loop-fusio // Reducer specialization can hide another producer/consumer pair behind this // compiler-generated call. Inline only that wrapper to expose further fusion, // charging the same per-caller growth budget used for producer substitution. +// For a producer mapped(source) that yields transform(value) for each source value: +// +// result = sum(mapped(gen(args)), start) +// +// becomes, after specializing the reducer and inlining its wrapper: +// +// source = gen(args) +// total = start +// for value in source: +// total = total + transform(value) +// result = total +// +// The inner gen/source pair is now visible to the loop-fusion rewrite below. void GeneratorLoopFusion::handle(CallInstr *call) { auto *parent = cast(getParentFunc()); auto *function = cast(util::getFunc(call->getCallee())); if (!parent || !function || function->getName() != "__sum_wrapper") return; + FusionVerifier verifier; verifier.process(function->getBody()); if (!verifier.valid || growth[parent->getId()] + verifier.nodes > 1024) return; + auto inlined = util::inlineCall(call, /*aggressive=*/true); if (!inlined) return; + for (auto *variable : inlined.newVars) parent->push_back(variable); + // The inliner implements early returns with a one-shot loop. Remember it so // later fusion accepts its targeted breaks and ignores it as a lifetime boundary. if (auto *expression = cast(inlined.result)) if (auto *body = cast(expression->getFlow())) if (auto *wrapper = cast(body->back())) wrappers.insert(wrapper->getId()); + growth[parent->getId()] += verifier.nodes; call->replaceAll(inlined.result); } +// Fuse a producer with its sole consumer, interleaving their ordinary code instead +// of suspending and resuming a generator. For example: +// +// def gen(limit): +// for item in range(limit): +// if stop(item): +// return finish() +// yield compute(item) +// iterator = gen(get_limit()) +// between() +// for value in iterator: +// consume(value) +// after() +// +// becomes (using labeled-break pseudocode): +// +// limit = get_limit() +// between() +// done: while True: +// for item in range(limit): +// if stop(item): +// finish() +// break done +// value = compute(item) +// consume(value) +// break +// after() +// +// Arguments are still evaluated at creation, but the producer body runs only at +// consumption. Its return ends iteration, not the caller. The consumer body is +// inserted unchanged, so a return there still exits the consuming function. void GeneratorLoopFusion::handle(ForFlow *loop) { if (loop->isParallel() || loop->isAsync()) return; + auto *parent = cast(getParentFunc()); if (!parent) return; + // Unpack a single-use tuple that forwards an iterator. Capture every element // at the original assignment, in order, even if only one feeds this loop; // dropping the other elements could discard effects. Then retry the direct case. + // + // packed = (gen(args), side_effect()) + // between() + // for value in packed[0]: consume(value) + // + // becomes: + // + // iterator = gen(args) + // unused = side_effect() + // between() + // for value in iterator: consume(value) + // + // This removes tuple forwarding, not side_effect(); iterator can now be fused. if (auto *extract = cast(loop->getIter())) { auto *variable = util::getVar(extract->getVal()); if (!variable || variable->isGlobal()) return; + IteratorUseVerifier uses(variable, loop, wrappers); uses.process(parent->getBody()); + auto *tuple = uses.valid() ? cast(uses.assignment->getRhs()) : nullptr; auto *constructor = tuple ? util::getFunc(tuple->getCallee()) : nullptr; auto *type = constructor ? cast(constructor->getParentType()) : nullptr; if (!type || type->getName() != "Tuple" || constructor->getUnmangledName() != Module::NEW_MAGIC_NAME) return; + auto index = cast(extract->getVal()->getType()) ->getMemberIndex(extract->getField()); if (index < 0 || index >= tuple->numArgs()) return; + auto *setup = loop->getModule()->Nr(); Var *selected = nullptr; int position = 0; @@ -413,11 +542,13 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { if (position++ == index) selected = variable; } + uses.assignment->replaceAll(setup); loop->setIter(loop->getModule()->Nr(selected)); handle(loop); return; } + // Accept either an immediate generator call or a uniquely owned local handle. // Retain the latter's assignment so argument evaluation stays at creation time. auto *call = cast(loop->getIter()); @@ -425,6 +556,7 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { if (auto *iterator = util::getVar(loop->getIter())) { if (iterator->isGlobal()) return; + IteratorUseVerifier uses(iterator, loop, wrappers); uses.process(parent->getBody()); if (uses.valid()) { @@ -432,6 +564,7 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { call = cast(creation->getRhs()); } } + auto *generator = call ? cast(util::getFunc(call->getCallee())) : nullptr; if (!generator || !generator->isGenerator() || !generator->getBody() || generator->isAsync() || parent == generator || @@ -439,10 +572,12 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { util::hasAttribute(generator, ast::getMangledFunc("std.internal.attributes", "noinline"))) return; + FusionVerifier verifier; verifier.process(generator->getBody()); ConsumerVerifier consumerVerifier(wrappers); consumerVerifier.process(loop->getBody()); + // One yield site lets us move, rather than duplicate, the consumer body. // The cumulative budget bounds growth when nested generators are fused in turn. if (!verifier.valid || verifier.yields != 1 || !consumerVerifier.valid || @@ -454,23 +589,29 @@ void GeneratorLoopFusion::handle(ForFlow *loop) { auto *setup = module->Nr(); std::unordered_map arguments; auto argument = generator->arg_begin(); + // Snapshot arguments once, left-to-right. For a named handle this setup replaces // its creation, not its use, preserving mutations and exceptions between the two. for (auto *value : *call) arguments.emplace((*argument++)->getId(), util::makeVar(value, setup, parent)); + if (creation) { creation->replaceAll(setup); setup = module->Nr(); } + util::CloneVisitor clone(module); auto *body = cast(clone.clone(generator->getBody(), parent, arguments)); + // A one-shot loop gives producer returns a common exit without returning from // the caller. Reaching the end also exits; only the producer's own loops repeat. auto *wrapper = module->Nr(); auto *exit = module->Nr(module->getBool(true), wrapper); wrappers.insert(exit->getId()); + FusionTransformer transformer(loop, exit); transformer.process(body); + wrapper->push_back(body); wrapper->push_back(module->Nr(exit)); setup->push_back(exit); From 461eb0992ca60b71e4e198a0a01b8fc2870a14e1 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 14:46:30 -0400 Subject: [PATCH 08/12] Fix plugin handling in LLVM codegen --- codon/cir/llvm/llvisitor.cpp | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/codon/cir/llvm/llvisitor.cpp b/codon/cir/llvm/llvisitor.cpp index bf2be42ff..5b4e2900b 100644 --- a/codon/cir/llvm/llvisitor.cpp +++ b/codon/cir/llvm/llvisitor.cpp @@ -535,7 +535,7 @@ void LLVMVisitor::writeToExecutable(const std::string &filename, if (plugins) { for (auto *plugin : *plugins) { - auto dylibPath = plugin->info.dylibPath; + const auto &dylibPath = plugin->info.getRuntimeDylibPath(); if (dylibPath.empty()) continue; @@ -556,7 +556,7 @@ void LLVMVisitor::writeToExecutable(const std::string &filename, if (plugins) { for (auto *plugin : *plugins) { if (plugin->info.linkArgs.empty()) { - auto dylibPath = plugin->info.dylibPath; + const auto &dylibPath = plugin->info.getRuntimeDylibPath(); if (dylibPath.empty()) continue; @@ -1203,6 +1203,10 @@ void LLVMVisitor::run(const std::vector &args, runLLVMPipeline(); Timer t1("llvm/jitlink"); + if (plugins) { + if (auto error = plugins->loadRuntimeLibraries()) + compilationError(llvm::toString(std::move(error))); + } for (auto &lib : libs) { std::string err; if (llvm::sys::DynamicLibrary::LoadLibraryPermanently(lib.c_str(), &err)) { From 07bb16a69505f7e0cad33a10597c4d2cf42f73cd Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 14:51:20 -0400 Subject: [PATCH 09/12] Bump version --- CMakeLists.txt | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index cedd9f765..6d9c25e52 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -1,10 +1,10 @@ cmake_minimum_required(VERSION 3.14) project( Codon - VERSION "0.20.2" + VERSION "0.20.3" HOMEPAGE_URL "https://github.com/exaloop/codon" DESCRIPTION "high-performance, extensible Python compiler") -set(CODON_JIT_PYTHON_VERSION "0.5.2") +set(CODON_JIT_PYTHON_VERSION "0.5.3") configure_file("${PROJECT_SOURCE_DIR}/cmake/config.h.in" "${PROJECT_SOURCE_DIR}/codon/config/config.h") configure_file("${PROJECT_SOURCE_DIR}/cmake/config.py.in" From eae4998a54bc831227788c8763d5581826b60d99 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 18:23:15 -0400 Subject: [PATCH 10/12] Fix sorting --- codon/runtime/numpy/sort.cpp | 132 ++++++++++++++++++++++++++++++---- stdlib/internal/sort.codon | 21 +++++- stdlib/numpy/sorting.codon | 2 +- test/numpy/test_io.codon | 7 +- test/numpy/test_sorting.codon | 19 +++++ test/stdlib/sort_test.codon | 28 ++++++-- 6 files changed, 181 insertions(+), 28 deletions(-) diff --git a/codon/runtime/numpy/sort.cpp b/codon/runtime/numpy/sort.cpp index 17ea5713c..2548913cd 100644 --- a/codon/runtime/numpy/sort.cpp +++ b/codon/runtime/numpy/sort.cpp @@ -1,34 +1,117 @@ // Copyright (C) 2022-2026 Exaloop Inc. #include "codon/runtime/lib.h" -#include "hwy/contrib/sort/vqsort-inl.h" + +#if (defined(__aarch64__) || defined(__arm64__)) && !defined(__ARM_FEATURE_SVE) +#define HWY_DISABLED_TARGETS HWY_ALL_SVE +#endif + +#include "hwy/contrib/sort/vqsort.h" #include +#include +#include #include +#undef HWY_TARGET_INCLUDE +#define HWY_TARGET_INCLUDE "codon/runtime/numpy/sort.cpp" +#include "hwy/contrib/sort/vqsort-inl.h" +#include "hwy/foreach_target.h" + +HWY_BEFORE_NAMESPACE(); +namespace { +namespace HWY_NAMESPACE { +namespace hn = hwy::HWY_NAMESPACE; + +template void sortPreservingZeros(T *data, size_t size) { + if (size < 2) + return; + const hn::CappedTag lanes; + const hn::RebindToUnsigned bits; + const hn::detail::MakeTraits traits; + const size_t width = hn::Lanes(lanes); + size_t negativeZeros = 0; + size_t nanCount = 0; + size_t index = 0; + for (; index + width <= size; index += width) { + const auto values = hn::LoadU(lanes, data + index); + const auto nans = hn::IsNaN(values); + const auto negatives = hn::Eq(hn::BitCast(bits, values), hn::SignBit(bits)); + negativeZeros += hn::CountTrue(bits, negatives); + nanCount += hn::CountTrue(lanes, nans); + const auto replace = hn::Or(nans, hn::RebindMask(lanes, negatives)); + hn::StoreU(hn::IfThenElse(replace, + hn::IfThenElse(nans, hn::Inf(lanes), hn::Zero(lanes)), + values), + lanes, data + index); + } + for (; index < size; ++index) { + const T value = data[index]; + if (value != value) { + ++nanCount; + data[index] = std::numeric_limits::infinity(); + } else if (value == 0) { + negativeZeros += std::signbit(value); + data[index] = 0; + } + } + HWY_ALIGN T buffer[hwy::SortConstants::BufBytes(HWY_MAX_BYTES) / sizeof(T)]; +#if VQSORT_ENABLED + if (!hn::detail::HandleSpecialCases(lanes, traits, data, size, buffer)) { + hn::detail::Recurse( + lanes, traits, data, size, buffer, hwy::detail::GetGeneratorStateStatic(), 50); + } +#else + hn::detail::HeapSort(traits, data, size); +#endif + if (negativeZeros) { + T *zeros = std::lower_bound(data, data + size - nanCount, T(0)); + hn::Fill(lanes, T(-0.0), negativeZeros, zeros); + } + if (nanCount) + hn::Fill(lanes, hn::GetLane(hn::NaN(lanes)), nanCount, data + size - nanCount); +} + +void sortListFloat32(float *data, size_t size) { sortPreservingZeros(data, size); } + +#ifndef __APPLE__ +void sortListFloat64(double *data, size_t size) { sortPreservingZeros(data, size); } +#endif +} // namespace HWY_NAMESPACE +} // namespace +HWY_AFTER_NAMESPACE(); + +#if HWY_ONCE namespace { +HWY_EXPORT(sortListFloat32); +#ifndef __APPLE__ +HWY_EXPORT(sortListFloat64); +#endif + template void sortPrimitive(T *data, int64_t size) { #ifdef __APPLE__ // macOS ARM64 benchmarks with 100M random/duplicate-heavy elements favored // libc++ for these 64-bit types and Highway for 16/32-bit types. - if constexpr (std::is_same_v || std::is_same_v || - std::is_same_v) { - if constexpr (std::is_same_v) { - // NaNs violate std::sort's default ordering; Highway places them last. - // This O(n) scan retains libc++'s specialized default-comparator path. - // Signed zeros compare equivalent and do not require the fallback. - // libc++ sort is significantly faster even with the initial scan. - if (std::any_of(data, data + size, [](T value) { return value != value; })) { - hwy::VQSort(data, size, hwy::SortAscending()); - return; - } - } + if constexpr (std::is_same_v || std::is_same_v) { std::sort(data, data + size); return; } #endif hwy::VQSort(data, size, hwy::SortAscending()); } + +template void sortNanLast(T *data, int64_t size) { + if (size < 2) + return; +#ifdef __APPLE__ + if constexpr (std::is_same_v) { + T *end = std::partition(data, data + size, [](T value) { return value == value; }); + std::sort(data, end); + return; + } +#endif + hwy::VQSort(data, size, hwy::SortAscending()); +} } // namespace SEQ_FUNC void cnp_sort_int16(int16_t *data, int64_t n) { sortPrimitive(data, n); } @@ -43,10 +126,29 @@ SEQ_FUNC void cnp_sort_int64(int64_t *data, int64_t n) { sortPrimitive(data, n); SEQ_FUNC void cnp_sort_uint64(uint64_t *data, int64_t n) { sortPrimitive(data, n); } -SEQ_FUNC void cnp_sort_float32(float *data, int64_t n) { sortPrimitive(data, n); } +SEQ_FUNC void cnp_sort_float32(float *data, int64_t n) { + if (n > 1) + HWY_DYNAMIC_DISPATCH(sortListFloat32)(data, n); +} + +SEQ_FUNC void cnp_sort_float64(double *data, int64_t n) { +#ifdef __APPLE__ + sortNanLast(data, n); +#else + if (n > 1) + HWY_DYNAMIC_DISPATCH(sortListFloat64)(data, n); +#endif +} + +SEQ_FUNC void cnp_sort_float32_nan_last(float *data, int64_t n) { + sortNanLast(data, n); +} -SEQ_FUNC void cnp_sort_float64(double *data, int64_t n) { sortPrimitive(data, n); } +SEQ_FUNC void cnp_sort_float64_nan_last(double *data, int64_t n) { + sortNanLast(data, n); +} SEQ_FUNC void cnp_sort_uint128(hwy::uint128_t *data, int64_t n) { hwy::VQSort(data, n, hwy::SortAscending()); } +#endif diff --git a/stdlib/internal/sort.codon b/stdlib/internal/sort.codon index 8c49f35e4..4696ba1a9 100644 --- a/stdlib/internal/sort.codon +++ b/stdlib/internal/sort.codon @@ -15,8 +15,17 @@ from C import cnp_sort_uint64(cobj, int) from C import cnp_sort_uint128(cobj, int) from C import cnp_sort_float32(cobj, int) from C import cnp_sort_float64(cobj, int) +from C import cnp_sort_float32_nan_last(cobj, int) +from C import cnp_sort_float64_nan_last(cobj, int) -def _try_native_sort(start: Ptr[T], n: int, T: type) -> bool: +def _try_native_sort(start: Ptr[T], n: int, nan_last: Literal[bool] = False, + T: type) -> bool: + """Try a native sort, returning False for unsupported types. + + For floats, the default List policy preserves signed-zero counts and leaves + NaN placement unspecified. nan_last selects NumPy's NaN-last policy, which + leaves zero signs unspecified. + """ if T is int or T is i64: cnp_sort_int64(start.as_byte(), n) elif T is i16: @@ -32,9 +41,15 @@ def _try_native_sort(start: Ptr[T], n: int, T: type) -> bool: elif T is u128: cnp_sort_uint128(start.as_byte(), n) elif T is float32: - cnp_sort_float32(start.as_byte(), n) + if nan_last: + cnp_sort_float32_nan_last(start.as_byte(), n) + else: + cnp_sort_float32(start.as_byte(), n) elif T is float: - cnp_sort_float64(start.as_byte(), n) + if nan_last: + cnp_sort_float64_nan_last(start.as_byte(), n) + else: + cnp_sort_float64(start.as_byte(), n) else: return False return True diff --git a/stdlib/numpy/sorting.codon b/stdlib/numpy/sorting.codon index 6b8c3bef6..c6d62a933 100644 --- a/stdlib/numpy/sorting.codon +++ b/stdlib/numpy/sorting.codon @@ -489,7 +489,7 @@ def _pdq_sort( leftmost = False def quicksort(start: Ptr[T], n: int, T: type): - if not _try_native_sort(start, n): + if not _try_native_sort(start, n, nan_last=True): _pdq_sort(start, 0, n, _floor_log2(n), True) def _apartial_insertion_sort(arr: Ptr[T], tosort: Ptr[int], begin: int, end: int, T: type): diff --git a/test/numpy/test_io.codon b/test/numpy/test_io.codon index d2221483c..6eb5eff9e 100644 --- a/test/numpy/test_io.codon +++ b/test/numpy/test_io.codon @@ -8,7 +8,7 @@ test_dir = "test/numpy/data" @test def test_array(): - a = np.array(np.empty((1, )), dtype=float) + a = np.array([1.25], dtype=float) f = test_dir + "/bin_0.npy" np.save(f, a) l = np.load(f, dtype=float) @@ -1708,10 +1708,11 @@ test_load_mismatched_dtype_v1() @test def test_load_empty_array_v1(): # Save an empty array and load - array = np.array(np.empty((1, )), dtype=int) + array = np.empty((0, ), dtype=int) filename = test_dir + "/binary_empty_array_v1.npy" np.save(filename, array) - loaded_array = np.load(filename) + loaded_array = np.load(filename, dtype=int) + assert loaded_array.shape == (0, ) assert np.array_equal(array, loaded_array) test_load_empty_array_v1() diff --git a/test/numpy/test_sorting.codon b/test/numpy/test_sorting.codon index 062ce92cd..fd00dcbfe 100644 --- a/test/numpy/test_sorting.codon +++ b/test/numpy/test_sorting.codon @@ -85,6 +85,25 @@ def test_sort_float_specials(dtype: type): value = float(data[index]) assert value == reference or (isnan(value) and isnan(reference)) + special = [nan, -0.0, 3.0, -nan, inf, 0.0, -2.0, -inf] + for size in (0, 1, 2, 7, 15, 16, 17, 63, 64, 65, 127, 128, 129, 1024): + values = [special[index % len(special)] for index in range(size)] + expected = sorted(value for value in values if not isnan(value)) + data = np.empty((size,), dtype=dtype) + for index, value in enumerate(values): + data[index] = value + inplace = data.copy() + inplace.sort() + for actual in (np.sort(data), inplace, np.sort(data[::-1])): + assert actual.shape == data.shape + for index, reference in enumerate(expected): + assert actual[index] == reference + for index in range(len(expected), size): + assert isnan(float(actual[index])) + for index, reference in enumerate(values): + value = float(data[index]) + assert value == reference or (isnan(value) and isnan(reference)) + def check_partitioned(vec, kth): if isinstance(kth, int): return check_partitioned(vec, (kth, )) diff --git a/test/stdlib/sort_test.codon b/test/stdlib/sort_test.codon index 7498e040b..02706afff 100644 --- a/test/stdlib/sort_test.codon +++ b/test/stdlib/sort_test.codon @@ -95,6 +95,7 @@ def test_standard_sort(): test_standard_sort() +@test def check_native_sort_type(T: type, lowest: T, highest: T): for size in (0, 1, 2, 7, 15, 16, 17, 63, 64, 65, 127, 128, 129, 1024): original = [T((index * 17) % 31 + 1) for index in range(size)] @@ -122,13 +123,27 @@ def check_native_sort_type(T: type, lowest: T, highest: T): else T((index * 17) % 31 + 1) for index in range(size)] +@test def check_special_float_sort(T: type): from math import copysign, isnan - for original in ([T(3), T(-0.0), T(0.0), T(1), T(-2)], - [T(-0.0), T(0.0), T(0.0), T(-0.0)], - [T("nan"), T("nan")], - [T("nan"), T(3), T(-1), T("nan"), T(2)], - [T(3), T("nan"), T(-1), T("inf"), T("-inf")]): + from random import Random + originals = [[T(3), T(-0.0), T(0.0), T(1), T(-2)], + [T(-0.0), T(0.0), T(0.0), T(-0.0)], + [T("nan"), T("nan")], + [T("nan"), T(3), T(-1), T("nan"), T(2)], + [T(3), T("nan"), T(-1), T("inf"), T("-inf")]] + mixed = [T("nan"), T(-0.0), T(3), T("inf"), T(0.0), T(-2), T("-inf")] + for size in (0, 1, 2, 7, 15, 16, 17, 63, 64, 65, 127, 128, 129, 1024, + 4097, 65536): + originals.append([T(-0.0) if index % 3 == 0 else T(0.0) + for index in range(size)]) + originals.append([mixed[index % len(mixed)] for index in range(size)]) + generator = Random(173) + stress = [T(generator.uniform(-1.0, 1.0)) for index in range(65536)] + for index in range(0, len(stress), 2): + stress[index] = mixed[generator.randrange(0, len(mixed))] + originals.append(stress) + for original in originals: for reverse in (False, True): expected = sorted((value for value in original if not isnan(float(value))), reverse=reverse, algorithm="pdq") @@ -139,7 +154,8 @@ def check_special_float_sort(T: type): assert sum(1 for value in actual if value == 0 and copysign(1.0, float(value)) < 0) == sum( 1 for value in original - if value == 0 and copysign(1.0, float(value)) < 0) + if value == 0 and copysign(1.0, float(value)) < 0), \ + f"{T.__name__}, size {len(original)}, reverse {reverse}: signed zeros changed" @test From 44cf5bfbaed30449e39c179b6455af7aa585fd01 Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 19:40:31 -0400 Subject: [PATCH 11/12] Fix Highway includes --- codon/runtime/numpy/sort.cpp | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/codon/runtime/numpy/sort.cpp b/codon/runtime/numpy/sort.cpp index 2548913cd..e90817aa4 100644 --- a/codon/runtime/numpy/sort.cpp +++ b/codon/runtime/numpy/sort.cpp @@ -13,10 +13,12 @@ #include #include +// clang-format off #undef HWY_TARGET_INCLUDE #define HWY_TARGET_INCLUDE "codon/runtime/numpy/sort.cpp" -#include "hwy/contrib/sort/vqsort-inl.h" #include "hwy/foreach_target.h" +#include "hwy/contrib/sort/vqsort-inl.h" +// clang-format on HWY_BEFORE_NAMESPACE(); namespace { From f5a7f9daf06ff64905c1ed66c16274eee6a3d96a Mon Sep 17 00:00:00 2001 From: "A. R. Shajii" Date: Sun, 27 Sep 2026 23:08:57 -0400 Subject: [PATCH 12/12] Fix test --- test/python/myextension.codon | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/python/myextension.codon b/test/python/myextension.codon index ee196f7f3..9091cf005 100644 --- a/test/python/myextension.codon +++ b/test/python/myextension.codon @@ -402,6 +402,6 @@ import numpy as np import numpy.pybridge def numpy_test(): - result = np.empty(2) + result = np.zeros(2) result[1:] = np.arange(1, 2) return result